diff --git a/src/.env.example b/src/.env.example index b5954540..ad1f3f54 100644 --- a/src/.env.example +++ b/src/.env.example @@ -69,3 +69,6 @@ ethereum_ws_url=wss://ethereum.publicnode.com # epay epay_pid= epay_key= +epay_default_token=usdt +epay_default_currency=cny +epay_default_network=tron diff --git a/src/config/config.go b/src/config/config.go index fa96b183..4b6b8570 100644 --- a/src/config/config.go +++ b/src/config/config.go @@ -182,6 +182,14 @@ func GetApiAuthToken() string { return viper.GetString("api_auth_token") } +func GetEpaySignKey() string { + key := strings.TrimSpace(GetEpayKey()) + if key != "" { + return key + } + return strings.TrimSpace(GetApiAuthToken()) +} + func GetRateApiUrl() string { rateURL := viper.GetString("api_rate_url") if rateURL == "" { @@ -348,3 +356,27 @@ func GetEpayPid() int { func GetEpayKey() string { return viper.GetString("epay_key") } + +func GetEpayDefaultToken() string { + token := strings.TrimSpace(viper.GetString("epay_default_token")) + if token == "" { + return "usdt" + } + return strings.ToLower(token) +} + +func GetEpayDefaultCurrency() string { + currency := strings.TrimSpace(viper.GetString("epay_default_currency")) + if currency == "" { + return "cny" + } + return strings.ToLower(currency) +} + +func GetEpayDefaultNetwork() string { + network := strings.TrimSpace(viper.GetString("epay_default_network")) + if network == "" { + return "tron" + } + return strings.ToLower(network) +} diff --git a/src/controller/comm/order_controller.go b/src/controller/comm/order_controller.go index 4c58c4d1..eae41ee4 100644 --- a/src/controller/comm/order_controller.go +++ b/src/controller/comm/order_controller.go @@ -1,10 +1,6 @@ package comm import ( - "encoding/json" - "fmt" - "log" - "github.com/assimon/luuu/model/request" "github.com/assimon/luuu/model/service" "github.com/assimon/luuu/util/constant" @@ -40,35 +36,22 @@ func (c *BaseCommController) SwitchNetwork(ctx echo.Context) (err error) { if err != nil { return c.FailJson(ctx, err) } - - jsonBytes, err := json.MarshalIndent(resp, "", " ") - if err != nil { - return c.FailJson(ctx, err) - } - - fmt.Printf("switch network response: \n%s", string(jsonBytes)) - return c.SucJson(ctx, resp) } func (c *BaseCommController) CreateTransactionAndRedirect(ctx echo.Context) (err error) { req := new(request.CreateTransactionRequest) if err = ctx.Bind(req); err != nil { - log.Println("bind request error:", err) return c.FailJson(ctx, constant.ParamsMarshalErr) } if err = c.ValidateStruct(ctx, req); err != nil { - log.Println("validate request error:", err) return c.FailJson(ctx, err) } resp, err := service.CreateTransaction(req) if err != nil { - log.Println("create transaction error:", err) return c.FailJson(ctx, err) } - fmt.Printf("create transaction response: %+v\n", resp) - tradeID := resp.TradeId ctx.Redirect(302, "/pay/checkout-counter/"+tradeID) diff --git a/src/controller/comm/pay_controller.go b/src/controller/comm/pay_controller.go index 66e68506..ec17c259 100644 --- a/src/controller/comm/pay_controller.go +++ b/src/controller/comm/pay_controller.go @@ -2,7 +2,6 @@ package comm import ( "encoding/json" - "fmt" "html/template" "net/http" "path/filepath" @@ -15,16 +14,33 @@ import ( // CheckoutCounter 收银台 func (c *BaseCommController) CheckoutCounter(ctx echo.Context) (err error) { + type pageData struct { + response.CheckoutCounterResponse + PaymentOptionsJSON template.JS + } + + buildPageData := func(resp response.CheckoutCounterResponse) pageData { + paymentOptionsJSON := template.JS("[]") + if len(resp.PaymentOptions) > 0 { + if b, err := json.Marshal(resp.PaymentOptions); err == nil { + paymentOptionsJSON = template.JS(string(b)) + } + } + return pageData{ + CheckoutCounterResponse: resp, + PaymentOptionsJSON: paymentOptionsJSON, + } + } + tradeId := ctx.Param("trade_id") resp, err := service.GetCheckoutCounterByTradeId(tradeId) if err != nil { - if err == service.ErrOrder { + if err == service.ErrOrderNotFound { tmpl, err := template.ParseFiles(filepath.Join(config.StaticFilePath, "index.html")) if err != nil { return ctx.String(http.StatusOK, err.Error()) } - emptyResp := response.CheckoutCounterResponse{} - return tmpl.Execute(ctx.Response(), emptyResp) + return tmpl.Execute(ctx.Response(), buildPageData(response.CheckoutCounterResponse{})) } return ctx.String(http.StatusOK, err.Error()) } @@ -33,13 +49,7 @@ func (c *BaseCommController) CheckoutCounter(ctx echo.Context) (err error) { return ctx.String(http.StatusOK, err.Error()) } - jsonByte, err := json.MarshalIndent(resp, "", " ") - if err != nil { - return ctx.String(http.StatusOK, err.Error()) - } - fmt.Printf("%v\n", string(jsonByte)) - - return tmpl.Execute(ctx.Response(), resp) + return tmpl.Execute(ctx.Response(), buildPageData(*resp)) } // CheckStatus 支付状态检测 diff --git a/src/controller/comm/supported_asset_controller.go b/src/controller/comm/supported_asset_controller.go index 72325d95..1b87d71d 100644 --- a/src/controller/comm/supported_asset_controller.go +++ b/src/controller/comm/supported_asset_controller.go @@ -1,13 +1,19 @@ package comm import ( + "fmt" "sort" "strconv" + "strings" + "github.com/assimon/luuu/config" "github.com/assimon/luuu/model/data" "github.com/assimon/luuu/model/response" + "github.com/assimon/luuu/model/service" "github.com/assimon/luuu/util/constant" + "github.com/assimon/luuu/util/walletaddr" "github.com/labstack/echo/v4" + "github.com/shopspring/decimal" ) type addSupportedAssetRequest struct { @@ -24,6 +30,19 @@ type updateSupportedAssetRequest struct { // GetSupportedAssets 对外公开可用链与 token 列表(无需鉴权,仅返回已启用项)。 func (c *BaseCommController) GetSupportedAssets(ctx echo.Context) error { + currency := strings.ToLower(strings.TrimSpace(ctx.QueryParam("currency"))) + amountText := strings.TrimSpace(ctx.QueryParam("amount")) + amountFilter := decimal.Zero + hasAmountFilter := false + if amountText != "" { + parsed, err := strconv.ParseFloat(amountText, 64) + if err != nil { + return c.FailJson(ctx, fmt.Errorf("invalid amount: %s", amountText)) + } + amountFilter = decimal.NewFromFloat(parsed) + hasAmountFilter = true + } + list, err := data.ListEnabledSupportedAssets() if err != nil { return c.FailJson(ctx, err) @@ -35,15 +54,37 @@ func (c *BaseCommController) GetSupportedAssets(ctx echo.Context) error { networkSet := make(map[string]struct{}) for _, w := range wallets { - networkSet[w.Network] = struct{}{} + network := walletaddr.NormalizeNetwork(w.Network) + address := walletaddr.Normalize(network, w.Address) + if !walletaddr.Validate(network, address) { + continue + } + networkSet[network] = struct{}{} } grouped := make(map[string][]string) for _, item := range list { - if _, ok := networkSet[item.Network]; !ok { + network := walletaddr.NormalizeNetwork(item.Network) + token := strings.ToUpper(strings.TrimSpace(item.Token)) + if _, ok := networkSet[network]; !ok { + continue + } + if currency != "" { + rate := config.GetRateForCoin(strings.ToLower(token), currency) + if rate <= 0 { + continue + } + if hasAmountFilter { + tokenAmount := amountFilter.Mul(decimal.NewFromFloat(rate)) + if tokenAmount.Cmp(decimal.NewFromFloat(service.UsdtMinimumPaymentAmount)) == -1 { + continue + } + } + } + if token == "" { continue } - grouped[item.Network] = append(grouped[item.Network], item.Token) + grouped[network] = append(grouped[network], token) } networks := make([]string, 0, len(grouped)) diff --git a/src/internal/testutil/testdb.go b/src/internal/testutil/testdb.go index 7c68ea2a..7e618db1 100644 --- a/src/internal/testutil/testdb.go +++ b/src/internal/testutil/testdb.go @@ -28,6 +28,8 @@ func SetupTestDatabases(t testing.TB) func() { viper.Set("queue_concurrency", 4) viper.Set("queue_poll_interval_ms", 50) viper.Set("api_auth_token", "test-token") + viper.Set("epay_key", "test-epay-key") + viper.Set("epay_pid", 1) config.HTTPAccessLog = false config.SQLDebug = false @@ -38,12 +40,21 @@ func SetupTestDatabases(t testing.TB) func() { mainDB := mustOpenSQLite(t, filepath.Join(t.TempDir(), "main.db")) runtimeDB := mustOpenSQLite(t, filepath.Join(t.TempDir(), "runtime.db")) - mustMigrate(t, mainDB, &mdb.Orders{}, &mdb.WalletAddress{}) + mustMigrate(t, mainDB, &mdb.Orders{}, &mdb.WalletAddress{}, &mdb.SupportedAsset{}) mustMigrate(t, runtimeDB, &mdb.TransactionLock{}) dao.Mdb = mainDB dao.RuntimeDB = runtimeDB + if err := mainDB.Create(&[]mdb.SupportedAsset{ + {Network: mdb.NetworkTron, Token: "USDT", Status: mdb.TokenStatusEnable}, + {Network: mdb.NetworkTron, Token: "TRX", Status: mdb.TokenStatusEnable}, + {Network: mdb.NetworkSolana, Token: "USDT", Status: mdb.TokenStatusEnable}, + {Network: mdb.NetworkEthereum, Token: "USDT", Status: mdb.TokenStatusEnable}, + }).Error; err != nil { + t.Fatalf("seed supported assets: %v", err) + } + return func() { closeDB(t, runtimeDB) closeDB(t, mainDB) diff --git a/src/model/data/order_data.go b/src/model/data/order_data.go index 430264b8..577a51e0 100644 --- a/src/model/data/order_data.go +++ b/src/model/data/order_data.go @@ -142,8 +142,7 @@ func GetSiblingSubOrders(parentTradeId string, excludeTradeId string) ([]mdb.Ord return orders, err } -// MarkParentOrderSuccess updates the parent order with the sub-order's payment details. -// Token and network are NOT overwritten — the parent keeps its original values. +// MarkParentOrderSuccess updates the parent order with the actual paid sub-order details. func MarkParentOrderSuccess(parentTradeId string, sub *mdb.Orders) (bool, error) { result := dao.Mdb.Model(&mdb.Orders{}). Where("trade_id = ?", parentTradeId). @@ -154,23 +153,67 @@ func MarkParentOrderSuccess(parentTradeId string, sub *mdb.Orders) (bool, error) "callback_confirm": mdb.CallBackConfirmNo, "actual_amount": sub.ActualAmount, "receive_address": sub.ReceiveAddress, + "token": sub.Token, + "network": sub.Network, }) return result.RowsAffected > 0, result.Error } -// MarkOrderSelected sets is_selected=true for the given trade_id. -func MarkOrderSelected(tradeId string) error { - return dao.Mdb.Model(&mdb.Orders{}). - Where("trade_id = ?", tradeId). - Update("is_selected", true).Error +func SetSelectedOrder(rootTradeId string, selectedTradeId string) error { + return dao.Mdb.Transaction(func(tx *gorm.DB) error { + if err := tx.Model(&mdb.Orders{}). + Where("trade_id = ? OR parent_trade_id = ?", rootTradeId, rootTradeId). + Update("is_selected", false).Error; err != nil { + return err + } + return tx.Model(&mdb.Orders{}). + Where("trade_id = ?", selectedTradeId). + Update("is_selected", true).Error + }) } -// RefreshOrderExpiration resets created_at to now so the expiration timer restarts. -// Called on the parent order when a sub-order is created or returned. -func RefreshOrderExpiration(tradeId string) error { - return dao.Mdb.Model(&mdb.Orders{}). - Where("trade_id = ?", tradeId). - Update("created_at", time.Now()).Error +func GetSelectedOrderInFamily(rootTradeId string) (*mdb.Orders, error) { + order := new(mdb.Orders) + err := dao.Mdb.Model(order). + Where("(trade_id = ? OR parent_trade_id = ?)", rootTradeId, rootTradeId). + Where("status = ?", mdb.StatusWaitPay). + Where("is_selected = ?", true). + Order("parent_trade_id asc, id asc"). + Limit(1). + Find(order).Error + return order, err +} + +func GetActiveOrdersInFamily(rootTradeId string) ([]mdb.Orders, error) { + var orders []mdb.Orders + err := dao.Mdb.Model(&mdb.Orders{}). + Where("(trade_id = ? OR parent_trade_id = ?)", rootTradeId, rootTradeId). + Where("status = ?", mdb.StatusWaitPay). + Find(&orders).Error + return orders, err +} + +func RefreshOrderFamilyExpiration(rootTradeId string, expirationTime time.Duration) error { + now := time.Now() + orders, err := GetActiveOrdersInFamily(rootTradeId) + if err != nil { + return err + } + if len(orders) == 0 { + return nil + } + tradeIDs := make([]string, 0, len(orders)) + for _, order := range orders { + tradeIDs = append(tradeIDs, order.TradeId) + } + if err = dao.Mdb.Model(&mdb.Orders{}). + Where("trade_id IN ?", tradeIDs). + Update("created_at", now).Error; err != nil { + return err + } + return dao.RuntimeDB.Model(&mdb.TransactionLock{}). + Where("trade_id IN ?", tradeIDs). + Update("expires_at", now.Add(expirationTime)).Error } // ResetCallbackConfirmOk sets callback_confirm back to Ok. diff --git a/src/model/data/supported_asset_data.go b/src/model/data/supported_asset_data.go index e691fba9..fde21927 100644 --- a/src/model/data/supported_asset_data.go +++ b/src/model/data/supported_asset_data.go @@ -121,3 +121,11 @@ func ListEnabledSupportedAssets() ([]mdb.SupportedAsset, error) { Find(&list).Error return list, err } + +func IsSupportedAssetEnabled(network, token string) (bool, error) { + asset, err := GetSupportedAssetByNetworkAndToken(network, token) + if err != nil { + return false, err + } + return asset.ID > 0 && asset.Status == mdb.TokenStatusEnable, nil +} diff --git a/src/model/data/wallet_address_data.go b/src/model/data/wallet_address_data.go index add0a83c..90b28a04 100644 --- a/src/model/data/wallet_address_data.go +++ b/src/model/data/wallet_address_data.go @@ -1,11 +1,10 @@ package data import ( - "strings" - "github.com/assimon/luuu/model/dao" "github.com/assimon/luuu/model/mdb" "github.com/assimon/luuu/util/constant" + "github.com/assimon/luuu/util/walletaddr" ) // AddWalletAddress 创建钱包 (默认 tron 网络,用于 Telegram 添加) @@ -13,23 +12,12 @@ func AddWalletAddress(address string) (*mdb.WalletAddress, error) { return AddWalletAddressWithNetwork(mdb.NetworkTron, address) } -// isEVMNetwork 判断是否是 EVM 网络 -func isEVMNetwork(network string) bool { - switch network { - case mdb.NetworkEthereum, mdb.NetworkBsc, mdb.NetworkPolygon, mdb.NetworkPlasma: - return true - } - return false -} - // AddWalletAddressWithNetwork 创建指定网络的钱包地址 func AddWalletAddressWithNetwork(network, address string) (*mdb.WalletAddress, error) { - network = strings.ToLower(strings.TrimSpace(network)) - address = strings.TrimSpace(address) - - // evm 网络地址统一小写,tron 和 solana 保持原样 - if isEVMNetwork(network) { - address = strings.ToLower(address) + network = walletaddr.NormalizeNetwork(network) + address = walletaddr.Normalize(network, address) + if !walletaddr.Validate(network, address) { + return nil, constant.InvalidWalletAddress } exist, err := GetWalletAddressByNetworkAndAddress(network, address) @@ -50,6 +38,8 @@ func AddWalletAddressWithNetwork(network, address string) (*mdb.WalletAddress, e // GetWalletAddressByNetworkAndAddress 通过网络和地址查询 func GetWalletAddressByNetworkAndAddress(network, address string) (*mdb.WalletAddress, error) { + network = walletaddr.NormalizeNetwork(network) + address = walletaddr.Normalize(network, address) walletAddress := new(mdb.WalletAddress) err := dao.Mdb.Model(walletAddress). Where("network = ?", network). @@ -87,6 +77,7 @@ func GetAvailableWalletAddress() ([]mdb.WalletAddress, error) { // GetAvailableWalletAddressByNetwork 获得指定网络的所有可用钱包地址 func GetAvailableWalletAddressByNetwork(network string) ([]mdb.WalletAddress, error) { + network = walletaddr.NormalizeNetwork(network) var list []mdb.WalletAddress err := dao.Mdb.Model(list). Where("status = ?", mdb.TokenStatusEnable). @@ -104,6 +95,7 @@ func GetAllWalletAddress() ([]mdb.WalletAddress, error) { // GetAllWalletAddressByNetwork 获得指定网络的所有钱包地址 func GetAllWalletAddressByNetwork(network string) ([]mdb.WalletAddress, error) { + network = walletaddr.NormalizeNetwork(network) var list []mdb.WalletAddress err := dao.Mdb.Model(list).Where("network = ?", network).Find(&list).Error return list, err diff --git a/src/model/mdb/orders_mdb.go b/src/model/mdb/orders_mdb.go index 9350b2c3..3606c786 100644 --- a/src/model/mdb/orders_mdb.go +++ b/src/model/mdb/orders_mdb.go @@ -31,6 +31,8 @@ type Orders struct { CallBackConfirm int `gorm:"column:callback_confirm;default:2" json:"callback_confirm"` IsSelected bool `gorm:"column:is_selected;default:false" json:"is_selected"` PaymentType string `gorm:"column:payment_type" json:"payment_type"` + PaymentChannel string `gorm:"column:payment_channel" json:"payment_channel"` + PaymentMerchantId string `gorm:"column:payment_merchant_id" json:"payment_merchant_id"` BaseModel } diff --git a/src/model/request/order_request.go b/src/model/request/order_request.go index b42e38e0..692ab086 100644 --- a/src/model/request/order_request.go +++ b/src/model/request/order_request.go @@ -4,16 +4,18 @@ import "github.com/gookit/validate" // CreateTransactionRequest 创建交易请求 type CreateTransactionRequest struct { - OrderId string `json:"order_id" validate:"required|maxLen:32"` - Currency string `json:"currency" validate:"required"` // 法币 如:cny - Token string `json:"token" validate:"required"` // 币种 如:usdt - Network string `json:"network" validate:"required"` // 网络 如:TRON - Amount float64 `json:"amount" validate:"required|isFloat|gt:0.01"` - NotifyUrl string `json:"notify_url" validate:"required"` - Signature string `json:"signature" validate:"required"` - RedirectUrl string `json:"redirect_url"` - Name string `json:"name"` - PaymentType string `json:"payment_type"` + OrderId string `json:"order_id" validate:"required|maxLen:32"` + Currency string `json:"currency" validate:"required"` // 法币 如:cny + Token string `json:"token" validate:"required"` // 币种 如:usdt + Network string `json:"network" validate:"required"` // 网络 如:TRON + Amount float64 `json:"amount" validate:"required|isFloat|min:0.01"` + NotifyUrl string `json:"notify_url" validate:"required"` + Signature string `json:"signature" validate:"required"` + RedirectUrl string `json:"redirect_url"` + Name string `json:"name"` + PaymentType string `json:"payment_type"` + PaymentChannel string `json:"payment_channel"` + PaymentMerchantId string `json:"payment_merchant_id"` } func (r CreateTransactionRequest) Translates() map[string]string { diff --git a/src/model/response/order_response.go b/src/model/response/order_response.go index 729ab7d5..68237e2b 100644 --- a/src/model/response/order_response.go +++ b/src/model/response/order_response.go @@ -21,6 +21,7 @@ type OrderNotifyResponse struct { ActualAmount float64 `json:"actual_amount"` // 订单实际需要支付的金额,保留4位小数 ReceiveAddress string `json:"receive_address"` // 收款钱包地址 Token string `json:"token"` // 所属币种 TRX USDT...... + Network string `json:"network"` // 所属网络 TRON ETH ... BlockTransactionId string `json:"block_transaction_id"` // 区块id Signature string `json:"signature"` // 签名 Status int `json:"status"` // 1:等待支付,2:支付成功,3:已过期 diff --git a/src/model/response/pay_response.go b/src/model/response/pay_response.go index 190194eb..8a953d11 100644 --- a/src/model/response/pay_response.go +++ b/src/model/response/pay_response.go @@ -1,17 +1,24 @@ package response +type CheckoutPaymentOption struct { + Token string `json:"token"` + Network string `json:"network"` +} + type CheckoutCounterResponse struct { - TradeId string `json:"trade_id"` // epusdt订单号 - Amount float64 `json:"amount"` // 订单金额,保留4位小数 法币金额 - ActualAmount float64 `json:"actual_amount"` // 订单实际需要支付的金额,保留4位小数 加密货币金额 - Token string `json:"token"` // 所属币种 TRX USDT...... - Currency string `json:"currency"` // 法币币种 CNY USD ... - ReceiveAddress string `json:"receive_address"` // 收款钱包地址 - Network string `json:"network"` // 网络 TRON ETH ... - ExpirationTime int64 `json:"expiration_time"` // 过期时间 时间戳 - RedirectUrl string `json:"redirect_url"` - CreatedAt int64 `json:"created_at"` // 订单创建时间 时间戳 - IsSelected bool `json:"is_selected"` + TradeId string `json:"trade_id"` // epusdt订单号 + Amount float64 `json:"amount"` // 订单金额,保留4位小数 法币金额 + ActualAmount float64 `json:"actual_amount"` // 订单实际需要支付的金额,保留4位小数 加密货币金额 + Token string `json:"token"` // 所属币种 TRX USDT...... + Currency string `json:"currency"` // 法币币种 CNY USD ... + ReceiveAddress string `json:"receive_address"` // 收款钱包地址 + Network string `json:"network"` // 网络 TRON ETH ... + Status int `json:"status"` // 订单状态 + ExpirationTime int64 `json:"expiration_time"` // 过期时间 时间戳 + RedirectUrl string `json:"redirect_url"` + CreatedAt int64 `json:"created_at"` // 订单创建时间 时间戳 + IsSelected bool `json:"is_selected"` + PaymentOptions []CheckoutPaymentOption `json:"payment_options,omitempty"` } type CheckStatusResponse struct { diff --git a/src/model/service/checkout_options.go b/src/model/service/checkout_options.go new file mode 100644 index 00000000..b982d747 --- /dev/null +++ b/src/model/service/checkout_options.go @@ -0,0 +1,145 @@ +package service + +import ( + "strings" + + "github.com/assimon/luuu/config" + "github.com/assimon/luuu/model/data" + "github.com/assimon/luuu/model/mdb" + "github.com/assimon/luuu/model/response" + "github.com/assimon/luuu/util/constant" + "github.com/shopspring/decimal" +) + +func ensurePaymentMethodAvailable(network, token string) error { + enabled, err := data.IsSupportedAssetEnabled(network, token) + if err != nil { + return err + } + if !enabled { + return constant.PaymentMethodUnavailable + } + return nil +} + +func buildCheckoutPaymentOptions(order *mdb.Orders) ([]response.CheckoutPaymentOption, error) { + rootTradeId := order.TradeId + if order.ParentTradeId != "" { + rootTradeId = order.ParentTradeId + } + + rootOrder, err := data.GetOrderInfoByTradeId(rootTradeId) + if err != nil { + return nil, err + } + if rootOrder.ID <= 0 { + return nil, constant.OrderNotExists + } + + activeSubOrders, err := data.GetActiveSubOrders(rootTradeId) + if err != nil { + return nil, err + } + + seen := map[string]struct{}{} + options := make([]response.CheckoutPaymentOption, 0, 1+len(activeSubOrders)) + appendOption := func(network, token string) { + network = strings.ToLower(strings.TrimSpace(network)) + token = strings.ToUpper(strings.TrimSpace(token)) + if network == "" || token == "" { + return + } + key := network + ":" + token + if _, ok := seen[key]; ok { + return + } + seen[key] = struct{}{} + options = append(options, response.CheckoutPaymentOption{ + Token: token, + Network: network, + }) + } + + appendOption(rootOrder.Network, rootOrder.Token) + for _, subOrder := range activeSubOrders { + appendOption(subOrder.Network, subOrder.Token) + } + + if len(activeSubOrders) >= MaxSubOrders { + return options, nil + } + + assets, err := data.ListEnabledSupportedAssets() + if err != nil { + return nil, err + } + wallets, err := data.GetAvailableWalletAddress() + if err != nil { + return nil, err + } + + networksWithWallet := make(map[string]struct{}, len(wallets)) + for _, wallet := range wallets { + network := strings.ToLower(strings.TrimSpace(wallet.Network)) + if network == "" { + continue + } + networksWithWallet[network] = struct{}{} + } + + for _, asset := range assets { + network := strings.ToLower(strings.TrimSpace(asset.Network)) + if _, ok := networksWithWallet[network]; !ok { + continue + } + rate := config.GetRateForCoin(strings.ToLower(asset.Token), strings.ToLower(rootOrder.Currency)) + if rate <= 0 { + continue + } + decimalTokenAmount := decimal.NewFromFloat(rootOrder.Amount).Mul(decimal.NewFromFloat(rate)) + if decimalTokenAmount.Cmp(decimal.NewFromFloat(UsdtMinimumPaymentAmount)) == -1 { + continue + } + appendOption(network, asset.Token) + } + + return options, nil +} + +func buildCheckoutResponse(order *mdb.Orders) (*response.CheckoutCounterResponse, error) { + if order.Status != mdb.StatusWaitPay { + return &response.CheckoutCounterResponse{ + TradeId: order.TradeId, + Amount: order.Amount, + ActualAmount: order.ActualAmount, + Token: order.Token, + Currency: order.Currency, + ReceiveAddress: order.ReceiveAddress, + Network: order.Network, + Status: order.Status, + ExpirationTime: order.CreatedAt.AddMinutes(config.GetOrderExpirationTime()).TimestampMilli(), + RedirectUrl: order.RedirectUrl, + CreatedAt: order.CreatedAt.TimestampMilli(), + IsSelected: order.IsSelected, + }, nil + } + paymentOptions, err := buildCheckoutPaymentOptions(order) + if err != nil { + return nil, err + } + return &response.CheckoutCounterResponse{ + TradeId: order.TradeId, + Amount: order.Amount, + ActualAmount: order.ActualAmount, + Token: order.Token, + Currency: order.Currency, + ReceiveAddress: order.ReceiveAddress, + Network: order.Network, + Status: order.Status, + ExpirationTime: order.CreatedAt.AddMinutes(config.GetOrderExpirationTime()).TimestampMilli(), + RedirectUrl: order.RedirectUrl, + CreatedAt: order.CreatedAt.TimestampMilli(), + IsSelected: order.IsSelected, + PaymentOptions: paymentOptions, + }, nil +} diff --git a/src/model/service/order_service.go b/src/model/service/order_service.go index 27aced12..5b904dff 100644 --- a/src/model/service/order_service.go +++ b/src/model/service/order_service.go @@ -63,6 +63,9 @@ func CreateTransaction(req *request.CreateTransactionRequest) (*response.CreateT if exist.ID > 0 { return nil, constant.OrderAlreadyExists } + if err = ensurePaymentMethodAvailable(network, token); err != nil { + return nil, err + } walletAddress, err := data.GetAvailableWalletAddressByNetwork(network) if err != nil { @@ -84,19 +87,21 @@ func CreateTransaction(req *request.CreateTransactionRequest) (*response.CreateT tx := dao.Mdb.Begin() order := &mdb.Orders{ - TradeId: tradeID, - OrderId: req.OrderId, - Amount: req.Amount, - Currency: currency, - ActualAmount: availableAmount, - ReceiveAddress: availableAddress, - Token: token, - Network: network, - Status: mdb.StatusWaitPay, - NotifyUrl: req.NotifyUrl, - RedirectUrl: req.RedirectUrl, - Name: req.Name, - PaymentType: req.PaymentType, + TradeId: tradeID, + OrderId: req.OrderId, + Amount: payAmount, + Currency: currency, + ActualAmount: availableAmount, + ReceiveAddress: availableAddress, + Token: token, + Network: network, + Status: mdb.StatusWaitPay, + NotifyUrl: req.NotifyUrl, + RedirectUrl: req.RedirectUrl, + Name: req.Name, + PaymentType: req.PaymentType, + PaymentChannel: strings.TrimSpace(req.PaymentChannel), + PaymentMerchantId: strings.TrimSpace(req.PaymentMerchantId), } if err = data.CreateOrderWithTransaction(tx, order); err != nil { tx.Rollback() @@ -305,9 +310,18 @@ func SwitchNetwork(req *request.SwitchNetworkRequest) (*response.CheckoutCounter // 2. Same token+network as parent → mark selected and return if strings.EqualFold(parent.Token, token) && strings.EqualFold(parent.Network, network) { - _ = data.MarkOrderSelected(parent.TradeId) + if err = data.SetSelectedOrder(parent.TradeId, parent.TradeId); err != nil { + return nil, err + } + if err = data.RefreshOrderFamilyExpiration(parent.TradeId, config.GetOrderExpirationTimeDuration()); err != nil { + return nil, err + } + parent, err = data.GetOrderInfoByTradeId(parent.TradeId) + if err != nil { + return nil, err + } parent.IsSelected = true - return buildCheckoutResponse(parent), nil + return buildCheckoutResponse(parent) } // 3. Existing active sub-order for this token+network → return it @@ -316,11 +330,18 @@ func SwitchNetwork(req *request.SwitchNetworkRequest) (*response.CheckoutCounter return nil, err } if existing.ID > 0 { - _ = data.MarkOrderSelected(parent.TradeId) - _ = data.MarkOrderSelected(existing.TradeId) - _ = data.RefreshOrderExpiration(parent.TradeId) + if err = data.SetSelectedOrder(parent.TradeId, existing.TradeId); err != nil { + return nil, err + } + if err = data.RefreshOrderFamilyExpiration(parent.TradeId, config.GetOrderExpirationTimeDuration()); err != nil { + return nil, err + } + existing, err = data.GetOrderInfoByTradeId(existing.TradeId) + if err != nil { + return nil, err + } existing.IsSelected = true - return buildCheckoutResponse(existing), nil + return buildCheckoutResponse(existing) } // 4. Check sub-order limit @@ -331,6 +352,9 @@ func SwitchNetwork(req *request.SwitchNetworkRequest) (*response.CheckoutCounter if count >= MaxSubOrders { return nil, constant.SubOrderLimitExceeded } + if err = ensurePaymentMethodAvailable(network, token); err != nil { + return nil, err + } // 5. Calculate amount for the new network rate := config.GetRateForCoin(strings.ToLower(token), strings.ToLower(parent.Currency)) @@ -365,22 +389,24 @@ func SwitchNetwork(req *request.SwitchNetworkRequest) (*response.CheckoutCounter // 7. Create sub-order tx := dao.Mdb.Begin() subOrder := &mdb.Orders{ - TradeId: subTradeID, - OrderId: subTradeID, // sub-order uses its own trade_id as order_id (unique constraint) - ParentTradeId: parent.TradeId, - Amount: parent.Amount, - Currency: parent.Currency, - ActualAmount: availableAmount, - ReceiveAddress: availableAddress, - Token: token, - Network: network, - Status: mdb.StatusWaitPay, - IsSelected: true, - NotifyUrl: "", - RedirectUrl: parent.RedirectUrl, - Name: parent.Name, - CallBackConfirm: mdb.CallBackConfirmOk, // don't trigger callback on sub-order - PaymentType: parent.PaymentType, + TradeId: subTradeID, + OrderId: subTradeID, // sub-order uses its own trade_id as order_id (unique constraint) + ParentTradeId: parent.TradeId, + Amount: parent.Amount, + Currency: parent.Currency, + ActualAmount: availableAmount, + ReceiveAddress: availableAddress, + Token: token, + Network: network, + Status: mdb.StatusWaitPay, + IsSelected: true, + NotifyUrl: "", + RedirectUrl: parent.RedirectUrl, + Name: parent.Name, + CallBackConfirm: mdb.CallBackConfirmOk, // don't trigger callback on sub-order + PaymentType: parent.PaymentType, + PaymentChannel: parent.PaymentChannel, + PaymentMerchantId: parent.PaymentMerchantId, } if err = data.CreateOrderWithTransaction(tx, subOrder); err != nil { tx.Rollback() @@ -394,24 +420,19 @@ func SwitchNetwork(req *request.SwitchNetworkRequest) (*response.CheckoutCounter } // Mark parent as selected and refresh its expiration to match the sub-order - _ = data.MarkOrderSelected(parent.TradeId) - _ = data.RefreshOrderExpiration(parent.TradeId) - - return buildCheckoutResponse(subOrder), nil -} - -func buildCheckoutResponse(order *mdb.Orders) *response.CheckoutCounterResponse { - return &response.CheckoutCounterResponse{ - TradeId: order.TradeId, - Amount: order.Amount, - ActualAmount: order.ActualAmount, - Token: order.Token, - Currency: order.Currency, - ReceiveAddress: order.ReceiveAddress, - Network: order.Network, - ExpirationTime: order.CreatedAt.AddMinutes(config.GetOrderExpirationTime()).TimestampMilli(), - RedirectUrl: order.RedirectUrl, - CreatedAt: order.CreatedAt.TimestampMilli(), - IsSelected: order.IsSelected, + if err = data.SetSelectedOrder(parent.TradeId, subOrder.TradeId); err != nil { + _ = data.UnLockTransactionByTradeId(subTradeID) + return nil, err } + if err = data.RefreshOrderFamilyExpiration(parent.TradeId, config.GetOrderExpirationTimeDuration()); err != nil { + _ = data.UnLockTransactionByTradeId(subTradeID) + return nil, err + } + subOrder, err = data.GetOrderInfoByTradeId(subTradeID) + if err != nil { + _ = data.UnLockTransactionByTradeId(subTradeID) + return nil, err + } + + return buildCheckoutResponse(subOrder) } diff --git a/src/model/service/order_service_test.go b/src/model/service/order_service_test.go index 99d73b37..719e55c8 100644 --- a/src/model/service/order_service_test.go +++ b/src/model/service/order_service_test.go @@ -13,6 +13,11 @@ import ( "github.com/assimon/luuu/util/constant" ) +const ( + testServiceTronAddr1 = "TLa2f6VPqDgRE67v1736s7bJ8Ray5wYjU7" + testServiceTronAddr2 = "TXLAQ63Xg1NAzckPwKHvzw7CSEmLMEqcdj" +) + func newCreateTransactionRequest(orderID string, amount float64) *request.CreateTransactionRequest { return &request.CreateTransactionRequest{ OrderId: orderID, @@ -28,7 +33,7 @@ func TestCreateTransactionAssignsIncrementedAmountsAndLocks(t *testing.T) { cleanup := testutil.SetupTestDatabases(t) defer cleanup() - if _, err := data.AddWalletAddress("wallet_1"); err != nil { + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { t.Fatalf("add wallet: %v", err) } @@ -47,7 +52,7 @@ func TestCreateTransactionAssignsIncrementedAmountsAndLocks(t *testing.T) { if got := fmt.Sprintf("%.2f", resp2.ActualAmount); got != "1.01" { t.Fatalf("second actual amount = %s, want 1.01", got) } - if resp1.ReceiveAddress != "wallet_1" || resp2.ReceiveAddress != "wallet_1" { + if resp1.ReceiveAddress != testServiceTronAddr1 || resp2.ReceiveAddress != testServiceTronAddr1 { t.Fatalf("unexpected receive addresses: %s, %s", resp1.ReceiveAddress, resp2.ReceiveAddress) } if resp1.Token != "USDT" || resp2.Token != "USDT" { @@ -75,7 +80,7 @@ func TestOrderProcessingMarksPaidAndReleasesLock(t *testing.T) { cleanup := testutil.SetupTestDatabases(t) defer cleanup() - if _, err := data.AddWalletAddress("wallet_1"); err != nil { + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { t.Fatalf("add wallet: %v", err) } @@ -123,7 +128,7 @@ func TestOrderProcessingRejectsDuplicateBlockForSameOrder(t *testing.T) { cleanup := testutil.SetupTestDatabases(t) defer cleanup() - if _, err := data.AddWalletAddress("wallet_1"); err != nil { + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { t.Fatalf("add wallet: %v", err) } @@ -165,7 +170,7 @@ func TestOrderProcessingDoesNotReviveExpiredOrder(t *testing.T) { cleanup := testutil.SetupTestDatabases(t) defer cleanup() - if _, err := data.AddWalletAddress("wallet_1"); err != nil { + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { t.Fatalf("add wallet: %v", err) } @@ -208,10 +213,10 @@ func TestOrderProcessingOnlyOneOrderClaimsABlockTransaction(t *testing.T) { cleanup := testutil.SetupTestDatabases(t) defer cleanup() - if _, err := data.AddWalletAddress("wallet_1"); err != nil { + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { t.Fatalf("add wallet: %v", err) } - if _, err := data.AddWalletAddress("wallet_2"); err != nil { + if _, err := data.AddWalletAddressWithNetwork(mdb.NetworkTron, testServiceTronAddr2); err != nil { t.Fatalf("add wallet: %v", err) } diff --git a/src/model/service/pay_service.go b/src/model/service/pay_service.go index 011707d1..167d8003 100644 --- a/src/model/service/pay_service.go +++ b/src/model/service/pay_service.go @@ -3,13 +3,12 @@ package service import ( "errors" - "github.com/assimon/luuu/config" "github.com/assimon/luuu/model/data" "github.com/assimon/luuu/model/mdb" "github.com/assimon/luuu/model/response" ) -var ErrOrder = errors.New("不存在待支付订单或已过期") +var ErrOrderNotFound = errors.New("订单不存在") // GetCheckoutCounterByTradeId returns checkout info for a pending order. func GetCheckoutCounterByTradeId(tradeId string) (*response.CheckoutCounterResponse, error) { @@ -17,22 +16,22 @@ func GetCheckoutCounterByTradeId(tradeId string) (*response.CheckoutCounterRespo if err != nil { return nil, err } - if orderInfo.ID <= 0 || orderInfo.Status != mdb.StatusWaitPay { - return nil, ErrOrder + if orderInfo.ID <= 0 { + return nil, ErrOrderNotFound } - - resp := &response.CheckoutCounterResponse{ - TradeId: orderInfo.TradeId, - Amount: orderInfo.Amount, - ActualAmount: orderInfo.ActualAmount, - Token: orderInfo.Token, - Currency: orderInfo.Currency, - ReceiveAddress: orderInfo.ReceiveAddress, - Network: orderInfo.Network, - ExpirationTime: orderInfo.CreatedAt.AddMinutes(config.GetOrderExpirationTime()).TimestampMilli(), - RedirectUrl: orderInfo.RedirectUrl, - CreatedAt: orderInfo.CreatedAt.TimestampMilli(), - IsSelected: orderInfo.IsSelected, + if orderInfo.Status != mdb.StatusWaitPay { + return buildCheckoutResponse(orderInfo) + } + rootTradeId := orderInfo.TradeId + if orderInfo.ParentTradeId != "" { + rootTradeId = orderInfo.ParentTradeId + } + selectedOrder, err := data.GetSelectedOrderInFamily(rootTradeId) + if err != nil { + return nil, err + } + if selectedOrder.ID > 0 { + orderInfo = selectedOrder } - return resp, nil + return buildCheckoutResponse(orderInfo) } diff --git a/src/model/service/pay_service_test.go b/src/model/service/pay_service_test.go new file mode 100644 index 00000000..a68b1c35 --- /dev/null +++ b/src/model/service/pay_service_test.go @@ -0,0 +1,266 @@ +package service + +import ( + "testing" + "time" + + "github.com/assimon/luuu/internal/testutil" + "github.com/assimon/luuu/model/dao" + "github.com/assimon/luuu/model/data" + "github.com/assimon/luuu/model/mdb" + "github.com/assimon/luuu/model/request" + "github.com/assimon/luuu/model/response" +) + +const testServiceSolAddr1 = "So11111111111111111111111111111111111111112" + +func hasPaymentOption(options []response.CheckoutPaymentOption, network, token string) bool { + for _, option := range options { + if option.Network == network && option.Token == token { + return true + } + } + return false +} + +func TestGetCheckoutCounterByTradeIdIncludesCurrentAndAvailableOptions(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { + t.Fatalf("add tron wallet: %v", err) + } + if _, err := data.AddWalletAddressWithNetwork(mdb.NetworkSolana, testServiceSolAddr1); err != nil { + t.Fatalf("add solana wallet: %v", err) + } + + resp, err := CreateTransaction(newCreateTransactionRequest("checkout_options", 1)) + if err != nil { + t.Fatalf("create transaction: %v", err) + } + + checkout, err := GetCheckoutCounterByTradeId(resp.TradeId) + if err != nil { + t.Fatalf("get checkout counter: %v", err) + } + if !hasPaymentOption(checkout.PaymentOptions, mdb.NetworkTron, "USDT") { + t.Fatalf("missing current tron option: %+v", checkout.PaymentOptions) + } + if !hasPaymentOption(checkout.PaymentOptions, mdb.NetworkSolana, "USDT") { + t.Fatalf("missing available solana option: %+v", checkout.PaymentOptions) + } +} + +func TestGetCheckoutCounterByTradeIdSkipsRateUnavailableOptions(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { + t.Fatalf("add tron wallet: %v", err) + } + if _, err := data.AddWalletAddressWithNetwork(mdb.NetworkSolana, testServiceSolAddr1); err != nil { + t.Fatalf("add solana wallet: %v", err) + } + + if err := dao.Mdb.Create(&mdb.SupportedAsset{ + Network: mdb.NetworkSolana, + Token: "USDC", + Status: mdb.TokenStatusEnable, + }).Error; err != nil { + t.Fatalf("create extra supported asset: %v", err) + } + + resp, err := CreateTransaction(newCreateTransactionRequest("checkout_rate_filter", 1)) + if err != nil { + t.Fatalf("create transaction: %v", err) + } + + checkout, err := GetCheckoutCounterByTradeId(resp.TradeId) + if err != nil { + t.Fatalf("get checkout counter: %v", err) + } + if hasPaymentOption(checkout.PaymentOptions, mdb.NetworkSolana, "USDC") { + t.Fatalf("unexpected rate-unavailable option: %+v", checkout.PaymentOptions) + } +} + +func TestGetCheckoutCounterByTradeIdKeepsCurrentOrderOptionAfterDisable(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { + t.Fatalf("add tron wallet: %v", err) + } + + resp, err := CreateTransaction(newCreateTransactionRequest("checkout_current_option", 1)) + if err != nil { + t.Fatalf("create transaction: %v", err) + } + + if err = dao.Mdb.Model(&mdb.SupportedAsset{}). + Where("network = ? AND token = ?", mdb.NetworkTron, "USDT"). + Update("status", mdb.TokenStatusDisable).Error; err != nil { + t.Fatalf("disable supported asset: %v", err) + } + + checkout, err := GetCheckoutCounterByTradeId(resp.TradeId) + if err != nil { + t.Fatalf("get checkout counter: %v", err) + } + if !hasPaymentOption(checkout.PaymentOptions, mdb.NetworkTron, "USDT") { + t.Fatalf("missing current order option after disable: %+v", checkout.PaymentOptions) + } +} + +func TestGetCheckoutCounterByTradeIdReturnsSelectedSubOrder(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { + t.Fatalf("add tron wallet: %v", err) + } + if _, err := data.AddWalletAddressWithNetwork(mdb.NetworkSolana, testServiceSolAddr1); err != nil { + t.Fatalf("add solana wallet: %v", err) + } + + resp, err := CreateTransaction(newCreateTransactionRequest("checkout_selected_sub", 1)) + if err != nil { + t.Fatalf("create transaction: %v", err) + } + + switchResp, err := SwitchNetwork(&request.SwitchNetworkRequest{ + TradeId: resp.TradeId, + Token: "usdt", + Network: "solana", + }) + if err != nil { + t.Fatalf("switch network: %v", err) + } + + checkout, err := GetCheckoutCounterByTradeId(resp.TradeId) + if err != nil { + t.Fatalf("get checkout counter: %v", err) + } + if checkout.TradeId != switchResp.TradeId { + t.Fatalf("expected selected sub-order %s, got %s", switchResp.TradeId, checkout.TradeId) + } + if checkout.Network != mdb.NetworkSolana || checkout.Token != "USDT" { + t.Fatalf("unexpected selected checkout payload: %+v", checkout) + } +} + +func TestSwitchNetworkRefreshesFamilyExpirationAndLocks(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { + t.Fatalf("add tron wallet: %v", err) + } + if _, err := data.AddWalletAddressWithNetwork(mdb.NetworkSolana, testServiceSolAddr1); err != nil { + t.Fatalf("add solana wallet: %v", err) + } + + resp, err := CreateTransaction(newCreateTransactionRequest("checkout_refresh_family", 1)) + if err != nil { + t.Fatalf("create transaction: %v", err) + } + firstSwitch, err := SwitchNetwork(&request.SwitchNetworkRequest{ + TradeId: resp.TradeId, + Token: "usdt", + Network: "solana", + }) + if err != nil { + t.Fatalf("switch network: %v", err) + } + + parentBefore, err := data.GetOrderInfoByTradeId(resp.TradeId) + if err != nil { + t.Fatalf("load parent order: %v", err) + } + subBefore, err := data.GetOrderInfoByTradeId(firstSwitch.TradeId) + if err != nil { + t.Fatalf("load sub-order: %v", err) + } + parentLockBefore := new(mdb.TransactionLock) + if err = dao.RuntimeDB.Where("trade_id = ?", resp.TradeId).Limit(1).Find(parentLockBefore).Error; err != nil { + t.Fatalf("load parent lock: %v", err) + } + subLockBefore := new(mdb.TransactionLock) + if err = dao.RuntimeDB.Where("trade_id = ?", firstSwitch.TradeId).Limit(1).Find(subLockBefore).Error; err != nil { + t.Fatalf("load sub lock: %v", err) + } + + time.Sleep(20 * time.Millisecond) + + secondSwitch, err := SwitchNetwork(&request.SwitchNetworkRequest{ + TradeId: resp.TradeId, + Token: "usdt", + Network: "solana", + }) + if err != nil { + t.Fatalf("switch network second time: %v", err) + } + if secondSwitch.TradeId != firstSwitch.TradeId { + t.Fatalf("expected existing sub-order %s, got %s", firstSwitch.TradeId, secondSwitch.TradeId) + } + + parentAfter, err := data.GetOrderInfoByTradeId(resp.TradeId) + if err != nil { + t.Fatalf("reload parent order: %v", err) + } + subAfter, err := data.GetOrderInfoByTradeId(firstSwitch.TradeId) + if err != nil { + t.Fatalf("reload sub-order: %v", err) + } + parentLockAfter := new(mdb.TransactionLock) + if err = dao.RuntimeDB.Where("trade_id = ?", resp.TradeId).Limit(1).Find(parentLockAfter).Error; err != nil { + t.Fatalf("reload parent lock: %v", err) + } + subLockAfter := new(mdb.TransactionLock) + if err = dao.RuntimeDB.Where("trade_id = ?", firstSwitch.TradeId).Limit(1).Find(subLockAfter).Error; err != nil { + t.Fatalf("reload sub lock: %v", err) + } + + if parentAfter.CreatedAt.TimestampMilli() <= parentBefore.CreatedAt.TimestampMilli() { + t.Fatalf("expected parent created_at to refresh: before=%v after=%v", parentBefore.CreatedAt, parentAfter.CreatedAt) + } + if subAfter.CreatedAt.TimestampMilli() <= subBefore.CreatedAt.TimestampMilli() { + t.Fatalf("expected sub-order created_at to refresh: before=%v after=%v", subBefore.CreatedAt, subAfter.CreatedAt) + } + if !parentLockAfter.ExpiresAt.After(parentLockBefore.ExpiresAt) { + t.Fatalf("expected parent lock to refresh: before=%v after=%v", parentLockBefore.ExpiresAt, parentLockAfter.ExpiresAt) + } + if !subLockAfter.ExpiresAt.After(subLockBefore.ExpiresAt) { + t.Fatalf("expected sub lock to refresh: before=%v after=%v", subLockBefore.ExpiresAt, subLockAfter.ExpiresAt) + } +} + +func TestGetCheckoutCounterByTradeIdReturnsTerminalOrder(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + if _, err := data.AddWalletAddress(testServiceTronAddr1); err != nil { + t.Fatalf("add tron wallet: %v", err) + } + + resp, err := CreateTransaction(newCreateTransactionRequest("checkout_terminal", 1)) + if err != nil { + t.Fatalf("create transaction: %v", err) + } + if err = dao.Mdb.Model(&mdb.Orders{}). + Where("trade_id = ?", resp.TradeId). + Update("status", mdb.StatusExpired).Error; err != nil { + t.Fatalf("expire order: %v", err) + } + + checkout, err := GetCheckoutCounterByTradeId(resp.TradeId) + if err != nil { + t.Fatalf("get checkout counter: %v", err) + } + if checkout.Status != mdb.StatusExpired { + t.Fatalf("expected expired checkout status, got %+v", checkout) + } + if len(checkout.PaymentOptions) != 0 { + t.Fatalf("expected no payment options for terminal order, got %+v", checkout.PaymentOptions) + } +} diff --git a/src/model/service/sol_task.go b/src/model/service/sol_task.go index 06731f94..120c8101 100644 --- a/src/model/service/sol_task.go +++ b/src/model/service/sol_task.go @@ -69,29 +69,18 @@ func SolCallBack(address string, wg *sync.WaitGroup) { log.Sugar.Errorf("[SOL][%s] failed to derive USDC ATA: %v", address, err) } + // 时间截止线:订单过期时间 + 5 分钟 + cutoffTime := time.Now().Add(-config.GetOrderExpirationTimeDuration() - 5*time.Minute).Unix() + // 拉取签名并去重 seen := make(map[string]bool) var result []solSignatureResult for _, queryAddr := range queryAddrs { - respBody, err := SolGetSignaturesForAddress(queryAddr, limit, "", "") - if err != nil { - log.Sugar.Errorf("[SOL][%s] SolGetSignaturesForAddress(%s) failed: %v", address, queryAddr, err) - continue - } - - resultBody := gjson.GetBytes(respBody, "result") - if !resultBody.Exists() || !resultBody.IsArray() { - log.Sugar.Errorf("[SOL][%s] unexpected response format for %s: %s", address, queryAddr, string(respBody)) - continue - } - - var batch []solSignatureResult - err = json.Unmarshal([]byte(resultBody.Raw), &batch) + batch, err := collectSolSignaturesForAddress(queryAddr, limit, cutoffTime) if err != nil { - log.Sugar.Errorf("[SOL][%s] failed to unmarshal signatures for %s: %v", address, queryAddr, err) + log.Sugar.Errorf("[SOL][%s] collect signatures for %s failed: %v", address, queryAddr, err) continue } - for _, sig := range batch { if !seen[sig.Signature] { seen[sig.Signature] = true @@ -116,9 +105,6 @@ func SolCallBack(address string, wg *sync.WaitGroup) { return *result[i].BlockTime > *result[j].BlockTime }) - // 时间截止线:订单过期时间 + 5 分钟 - cutoffTime := time.Now().Add(-config.GetOrderExpirationTimeDuration() - 5*time.Minute).Unix() - log.Sugar.Debugf("[SOL][%s] fetched %d unique signatures from %d addresses, cutoff=%d", address, len(result), len(queryAddrs), cutoffTime) @@ -256,6 +242,51 @@ type solSignatureResult struct { BlockTime *int64 `json:"blockTime"` } +func collectSolSignaturesForAddress(address string, limit int, cutoffTime int64) ([]solSignatureResult, error) { + return collectSolSignaturesForAddressWithFetcher(SolGetSignaturesForAddress, address, limit, cutoffTime) +} + +func collectSolSignaturesForAddressWithFetcher(fetcher func(string, int, string, string) ([]byte, error), address string, limit int, cutoffTime int64) ([]solSignatureResult, error) { + beforeSig := "" + result := make([]solSignatureResult, 0, limit) + for { + respBody, err := fetcher(address, limit, "", beforeSig) + if err != nil { + return nil, err + } + + resultBody := gjson.GetBytes(respBody, "result") + if !resultBody.Exists() || !resultBody.IsArray() { + return nil, fmt.Errorf("unexpected response format for %s: %s", address, string(respBody)) + } + + var batch []solSignatureResult + if err = json.Unmarshal([]byte(resultBody.Raw), &batch); err != nil { + return nil, err + } + if len(batch) == 0 { + return result, nil + } + + stop := false + for _, sig := range batch { + if sig.BlockTime != nil && *sig.BlockTime < cutoffTime { + stop = true + break + } + result = append(result, sig) + } + if stop || len(batch) < limit { + return result, nil + } + + beforeSig = batch[len(batch)-1].Signature + if beforeSig == "" { + return result, nil + } + } +} + // SolRetryClient 发送 Solana JSON-RPC 请求,自动重试 func SolRetryClient(method string, params []interface{}) ([]byte, error) { client := resty.New() @@ -437,7 +468,6 @@ func ADJustAmount(amount uint64, decimals int) float64 { func MatchUsdtAtaAddress(address string, ataTo string) bool { ata, err := FindATAAddress(address, USDT_Mint) if err != nil { - fmt.Printf("FindATAAddress failed: %v\n", err) return false } @@ -447,7 +477,6 @@ func MatchUsdtAtaAddress(address string, ataTo string) bool { func MatchUsdcAtaAddress(address string, ataTo string) bool { ata, err := FindATAAddress(address, USDC_Mint) if err != nil { - fmt.Printf("FindATAAddress failed: %v\n", err) return false } @@ -457,7 +486,6 @@ func MatchUsdcAtaAddress(address string, ataTo string) bool { func MatchAtaAddress(address string, mint string, ataTo string) bool { ata, err := FindATAAddress(address, mint) if err != nil { - fmt.Printf("FindATAAddress failed: %v\n", err) return false } diff --git a/src/model/service/sol_task_test.go b/src/model/service/sol_task_test.go index 924a2a6f..171ef06a 100644 --- a/src/model/service/sol_task_test.go +++ b/src/model/service/sol_task_test.go @@ -2,17 +2,26 @@ package service import ( "encoding/json" - "fmt" "testing" "github.com/tidwall/gjson" ) -func TestSolClientHealthy(t *testing.T) { - bodyData, err := SolRetryClient("getHealth", nil) +func skipOnSolRPCError(t *testing.T, bodyData []byte, err error) { + t.Helper() + if err != nil { - t.Fatalf("SolRetryClient failed: %v", err) + t.Skipf("skip due to solana rpc request error: %v", err) + } + + if rpcErr := gjson.GetBytes(bodyData, "error"); rpcErr.Exists() { + t.Skipf("skip due to solana rpc error: %s", rpcErr.Raw) } +} + +func TestSolClientHealthy(t *testing.T) { + bodyData, err := SolRetryClient("getHealth", nil) + skipOnSolRPCError(t, bodyData, err) var result map[string]interface{} err = json.Unmarshal(bodyData, &result) @@ -37,9 +46,7 @@ func TestSolClientGetSignaturesForAddress(t *testing.T) { address := "2uFTf9TZ8gd7Kg6hkb79TxfaeNpaAgpJ8uVHguv2Yweu" bodyData, err := SolRetryClient("getSignaturesForAddress", []interface{}{address, map[string]interface{}{"commitment": "finalized", "limit": 100}}) - if err != nil { - t.Fatalf("SolRetryClient failed: %v", err) - } + skipOnSolRPCError(t, bodyData, err) var result map[string]interface{} err = json.Unmarshal(bodyData, &result) @@ -61,10 +68,7 @@ func TestSolClientGetTransaction(t *testing.T) { sig := "2aEoNykk4ZJ27C3y7EDJiQUc7GFnnsMe7ofFzB73swGL8kTxSBFCnwzWw3jzr3BND7k8hx15fZHUUAbG1XemNFe5" txData, err := SolRetryClient("getTransaction", []interface{}{sig, map[string]interface{}{"encoding": "jsonParsed", "commitment": "finalized"}}) - if err != nil { - t.Fatalf("SolRetryClient failed: %v", err) - } - fmt.Printf("%v\n", string(txData)) + skipOnSolRPCError(t, txData, err) var result map[string]interface{} err = json.Unmarshal(txData, &result) @@ -113,6 +117,40 @@ func TestFindATAAddress(t *testing.T) { } } +func TestCollectSolSignaturesForAddressWithFetcherPaginates(t *testing.T) { + var calls int + results, err := collectSolSignaturesForAddressWithFetcher(func(address string, limit int, untilSig string, beforeSig string) ([]byte, error) { + calls++ + switch calls { + case 1: + if beforeSig != "" { + t.Fatalf("unexpected first before signature: %q", beforeSig) + } + return []byte(`{"result":[{"signature":"sig-1","slot":1,"err":null,"blockTime":200},{"signature":"sig-2","slot":2,"err":null,"blockTime":190}]}`), nil + case 2: + if beforeSig != "sig-2" { + t.Fatalf("unexpected second before signature: %q", beforeSig) + } + return []byte(`{"result":[{"signature":"sig-3","slot":3,"err":null,"blockTime":180},{"signature":"sig-old","slot":4,"err":null,"blockTime":90}]}`), nil + default: + t.Fatalf("unexpected extra fetch call %d", calls) + return nil, nil + } + }, "test-address", 2, 100) + if err != nil { + t.Fatalf("collectSolSignaturesForAddressWithFetcher failed: %v", err) + } + if calls != 2 { + t.Fatalf("expected 2 fetch calls, got %d", calls) + } + if len(results) != 3 { + t.Fatalf("expected 3 signatures before cutoff, got %+v", results) + } + if results[0].Signature != "sig-1" || results[1].Signature != "sig-2" || results[2].Signature != "sig-3" { + t.Fatalf("unexpected paginated signatures: %+v", results) + } +} + func TestMatchATAAddress(t *testing.T) { owner := "2uFTf9TZ8gd7Kg6hkb79TxfaeNpaAgpJ8uVHguv2Yweu" mint := "4k3Dyjzvzp8eMZWUXbBCjEvwSkkk59S5iCNLY3QrkX6R" // ray token @@ -226,9 +264,7 @@ func TestParseTransferInfoFromInstruction_SplTransfer(t *testing.T) { // SPL Token "transfer" (no mint in instruction, must look up from postTokenBalances) sig := "3tZTwLrvmiZ59h4UzyMHPd7DPux7t9eXZgkUvEfquaoSuERrPSRNzWuSHKQM2fbiCWFDGNqoLpu2kLZnfoegVpqN" txData, err := SolGetTransaction(sig) - if err != nil { - t.Fatalf("SolGetTransaction failed: %v", err) - } + skipOnSolRPCError(t, txData, err) instructions := gjson.GetBytes(txData, "result.transaction.message.instructions").Array() var found bool @@ -268,9 +304,7 @@ func TestParseTransferInfoFromInstruction_TransferChecked(t *testing.T) { // SPL Token "transferChecked" (has mint and tokenAmount in instruction) sig := "2aEoNykk4ZJ27C3y7EDJiQUc7GFnnsMe7ofFzB73swGL8kTxSBFCnwzWw3jzr3BND7k8hx15fZHUUAbG1XemNFe5" txData, err := SolGetTransaction(sig) - if err != nil { - t.Fatalf("SolGetTransaction failed: %v", err) - } + skipOnSolRPCError(t, txData, err) instructions := gjson.GetBytes(txData, "result.transaction.message.instructions").Array() var found bool @@ -313,9 +347,7 @@ func TestParseTransferInfoFromInstruction_SystemTransfer(t *testing.T) { // System program SOL transfer sig := "5pNMonUBvLVpxXTmyd5CGVBs49W6781g2ACnrCXhbmtz58KENYA7HSqu6hQkQweg3qQboRd8WAscphNAtiq9UtZZ" txData, err := SolGetTransaction(sig) - if err != nil { - t.Fatalf("SolGetTransaction failed: %v", err) - } + skipOnSolRPCError(t, txData, err) instructions := gjson.GetBytes(txData, "result.transaction.message.instructions").Array() transferCount := 0 diff --git a/src/model/service/task_service.go b/src/model/service/task_service.go index 131da67d..09acd887 100644 --- a/src/model/service/task_service.go +++ b/src/model/service/task_service.go @@ -8,6 +8,7 @@ import ( "net/http" "strings" "sync" + "time" "github.com/assimon/luuu/config" tron "github.com/assimon/luuu/crypto" @@ -27,6 +28,7 @@ import ( ) const TRC20_USDT_ID = "TR7NHqjeKQxGTCi8q8ZY4pL8otSzgjLj6t" +const tronGridPageLimit = 100 func Trc20CallBack(address string, wg *sync.WaitGroup) { defer wg.Done() @@ -43,6 +45,42 @@ func Trc20CallBack(address string, wg *sync.WaitGroup) { innerWg.Wait() } +func cloneQueryParams(src map[string]string) map[string]string { + dst := make(map[string]string, len(src)+1) + for k, v := range src { + dst[k] = v + } + return dst +} + +func walkTronGridPages(fetchPage func(map[string]string) ([]byte, error), baseQuery map[string]string, visit func(gjson.Result) bool) error { + fingerprint := "" + for { + query := cloneQueryParams(baseQuery) + if fingerprint != "" { + query["fingerprint"] = fingerprint + } + body, err := fetchPage(query) + if err != nil { + return err + } + records := gjson.GetBytes(body, "data").Array() + if len(records) == 0 { + return nil + } + for _, record := range records { + if !visit(record) { + return nil + } + } + next := gjson.GetBytes(body, "meta.fingerprint").String() + if next == "" { + return nil + } + fingerprint = next + } +} + func checkTrxTransfers(address string, wg *sync.WaitGroup) { defer wg.Done() defer func() { @@ -52,63 +90,58 @@ func checkTrxTransfers(address string, wg *sync.WaitGroup) { }() client := http_client.GetHttpClient() - startTime := carbon.Now().AddHours(-24).TimestampMilli() + startTime := time.Now().Add(-config.GetOrderExpirationTimeDuration() - 5*time.Minute).UnixMilli() endTime := carbon.Now().TimestampMilli() url := fmt.Sprintf("https://api.trongrid.io/v1/accounts/%s/transactions", address) - - resp, err := client.R().SetQueryParams(map[string]string{ + err := walkTronGridPages(func(query map[string]string) ([]byte, error) { + resp, err := client.R().SetQueryParams(query).SetHeader("TRON-PRO-API-KEY", config.TRON_GRID_API_KEY).Get(url) + if err != nil { + return nil, err + } + if resp.StatusCode() != http.StatusOK { + return nil, fmt.Errorf("TRX API returned status %d", resp.StatusCode()) + } + if !gjson.GetBytes(resp.Body(), "success").Bool() { + return nil, errors.New("TRX API response indicates failure") + } + return resp.Body(), nil + }, map[string]string{ "order_by": "block_timestamp,desc", - "limit": "100", + "limit": stdutil.ToString(tronGridPageLimit), "only_to": "true", "min_timestamp": stdutil.ToString(startTime), "max_timestamp": stdutil.ToString(endTime), - }).SetHeader("TRON-PRO-API-KEY", config.TRON_GRID_API_KEY).Get(url) - if err != nil { - panic(err) - } - if resp.StatusCode() != http.StatusOK { - panic(fmt.Sprintf("TRX API returned status %d", resp.StatusCode())) - } - - success := gjson.GetBytes(resp.Body(), "success").Bool() - if !success { - panic("TRX API response indicates failure") - } - - transfers := gjson.GetBytes(resp.Body(), "data").Array() - if len(transfers) == 0 { - log.Sugar.Debugf("[TRX][%s] no transfer records found", address) - return - } - log.Sugar.Debugf("[TRX][%s] fetched %d transfer records", address, len(transfers)) - - for i, transfer := range transfers { + }, func(transfer gjson.Result) bool { + blockTimestamp := transfer.Get("block_timestamp").Int() + if blockTimestamp < startTime { + return false + } if transfer.Get("raw_data.contract.0.type").String() != "TransferContract" { - continue + return true } if transfer.Get("ret.0.contractRet").String() != "SUCCESS" { - continue + return true } toAddressHex := transfer.Get("raw_data.contract.0.parameter.value.to_address").String() toBytes, err := hex.DecodeString(toAddressHex) if err != nil { - log.Sugar.Errorf("[TRX][%s] decode address failed on tx #%d: %v", address, i, err) - continue + log.Sugar.Errorf("[TRX][%s] decode address failed: %v", address, err) + return true } if tron.EncodeCheck(toBytes) != address { - continue + return true } rawAmount := transfer.Get("raw_data.contract.0.parameter.value.amount").String() decimalQuant, err := decimal.NewFromString(rawAmount) if err != nil { - log.Sugar.Errorf("[TRX][%s] parse amount failed on tx #%d: %v", address, i, err) - continue + log.Sugar.Errorf("[TRX][%s] parse amount failed: %v", address, err) + return true } amount := math.MustParsePrecFloat64(decimalQuant.Div(decimal.NewFromInt(1000000)).InexactFloat64(), 2) if amount <= 0 { - continue + return true } txID := transfer.Get("txID").String() @@ -118,7 +151,7 @@ func checkTrxTransfers(address string, wg *sync.WaitGroup) { } if tradeID == "" { log.Sugar.Debugf("[TRX][%s] skip unmatched tx hash=%s amount=%.2f", address, txID, amount) - continue + return true } log.Sugar.Infof("[TRX][%s] matched trade_id=%s hash=%s amount=%.2f", address, tradeID, txID, amount) @@ -126,11 +159,10 @@ func checkTrxTransfers(address string, wg *sync.WaitGroup) { if err != nil { panic(err) } - blockTimestamp := transfer.Get("block_timestamp").Int() createTime := order.CreatedAt.TimestampMilli() if blockTimestamp < createTime { log.Sugar.Warnf("[TRX][%s] skip tx %s because block time %d is before order create time %d", address, txID, blockTimestamp, createTime) - continue + return true } req := &request.OrderProcessingRequest{ @@ -145,13 +177,17 @@ func checkTrxTransfers(address string, wg *sync.WaitGroup) { if err != nil { if errors.Is(err, constant.OrderBlockAlreadyProcess) || errors.Is(err, constant.OrderStatusConflict) { log.Sugar.Infof("[TRX][%s] skip resolved transfer trade_id=%s hash=%s err=%v", address, tradeID, txID, err) - continue + return true } panic(err) } sendPaymentNotification(order) log.Sugar.Infof("[TRX][%s] payment processed trade_id=%s hash=%s", address, tradeID, txID) + return true + }) + if err != nil { + panic(err) } } @@ -164,54 +200,49 @@ func checkTrc20Transfers(address string, wg *sync.WaitGroup) { }() client := http_client.GetHttpClient() - startTime := carbon.Now().AddHours(-24).TimestampMilli() + startTime := time.Now().Add(-config.GetOrderExpirationTimeDuration() - 5*time.Minute).UnixMilli() endTime := carbon.Now().TimestampMilli() url := fmt.Sprintf("https://api.trongrid.io/v1/accounts/%s/transactions/trc20", address) - - resp, err := client.R().SetQueryParams(map[string]string{ + err := walkTronGridPages(func(query map[string]string) ([]byte, error) { + resp, err := client.R().SetQueryParams(query).SetHeader("TRON-PRO-API-KEY", config.TRON_GRID_API_KEY).Get(url) + if err != nil { + return nil, err + } + if resp.StatusCode() != http.StatusOK { + return nil, fmt.Errorf("TRC20 API returned status %d", resp.StatusCode()) + } + if !gjson.GetBytes(resp.Body(), "success").Bool() { + return nil, errors.New("TRC20 API response indicates failure") + } + return resp.Body(), nil + }, map[string]string{ "order_by": "block_timestamp,desc", - "limit": "100", + "limit": stdutil.ToString(tronGridPageLimit), "only_to": "true", "min_timestamp": stdutil.ToString(startTime), "max_timestamp": stdutil.ToString(endTime), - }).SetHeader("TRON-PRO-API-KEY", config.TRON_GRID_API_KEY).Get(url) - if err != nil { - panic(err) - } - if resp.StatusCode() != http.StatusOK { - panic(fmt.Sprintf("TRC20 API returned status %d", resp.StatusCode())) - } - - success := gjson.GetBytes(resp.Body(), "success").Bool() - if !success { - panic("TRC20 API response indicates failure") - } - - transfers := gjson.GetBytes(resp.Body(), "data").Array() - if len(transfers) == 0 { - log.Sugar.Debugf("[TRC20][%s] no transfer records found", address) - return - } - log.Sugar.Debugf("[TRC20][%s] fetched %d transfer records", address, len(transfers)) - - for i, transfer := range transfers { + }, func(transfer gjson.Result) bool { + blockTimestamp := transfer.Get("block_timestamp").Int() + if blockTimestamp < startTime { + return false + } if transfer.Get("token_info.address").String() != TRC20_USDT_ID { - continue + return true } if transfer.Get("to").String() != address { - continue + return true } valueStr := transfer.Get("value").String() decimalQuant, err := decimal.NewFromString(valueStr) if err != nil { - log.Sugar.Errorf("[TRC20][%s] parse value failed on tx #%d: %v", address, i, err) - continue + log.Sugar.Errorf("[TRC20][%s] parse value failed: %v", address, err) + return true } tokenDecimals := transfer.Get("token_info.decimals").Int() amount := math.MustParsePrecFloat64(decimalQuant.Div(decimal.New(1, int32(tokenDecimals))).InexactFloat64(), 2) if amount <= 0 { - continue + return true } txID := transfer.Get("transaction_id").String() @@ -221,7 +252,7 @@ func checkTrc20Transfers(address string, wg *sync.WaitGroup) { } if tradeID == "" { log.Sugar.Debugf("[TRC20][%s] skip unmatched tx hash=%s amount=%.2f", address, txID, amount) - continue + return true } log.Sugar.Infof("[TRC20][%s] matched trade_id=%s hash=%s amount=%.2f", address, tradeID, txID, amount) @@ -229,11 +260,10 @@ func checkTrc20Transfers(address string, wg *sync.WaitGroup) { if err != nil { panic(err) } - blockTimestamp := transfer.Get("block_timestamp").Int() createTime := order.CreatedAt.TimestampMilli() if blockTimestamp < createTime { log.Sugar.Warnf("[TRC20][%s] skip tx %s because block time %d is before order create time %d", address, txID, blockTimestamp, createTime) - continue + return true } req := &request.OrderProcessingRequest{ @@ -248,13 +278,17 @@ func checkTrc20Transfers(address string, wg *sync.WaitGroup) { if err != nil { if errors.Is(err, constant.OrderBlockAlreadyProcess) || errors.Is(err, constant.OrderStatusConflict) { log.Sugar.Infof("[TRC20][%s] skip resolved transfer trade_id=%s hash=%s err=%v", address, tradeID, txID, err) - continue + return true } panic(err) } sendPaymentNotification(order) log.Sugar.Infof("[TRC20][%s] payment processed trade_id=%s hash=%s", address, tradeID, txID) + return true + }) + if err != nil { + panic(err) } } @@ -343,6 +377,11 @@ func TryProcessEvmERC20Transfer(chainNetwork string, contract common.Address, to log.Sugar.Warnf("[%s-%s][%s] load order: %v", net, tokenSym, walletAddr, err) return } + if blockTsMs > 0 && blockTsMs < order.CreatedAt.TimestampMilli() { + log.Sugar.Warnf("[%s-%s][%s] skip tx hash=%s because block time %d is before order create time %d", + net, tokenSym, walletAddr, txHash, blockTsMs, order.CreatedAt.TimestampMilli()) + return + } if strings.ToLower(strings.TrimSpace(order.Network)) != chainNetwork { log.Sugar.Warnf("[%s-%s][%s] skip trade_id=%s network=%q", net, tokenSym, walletAddr, tradeID, order.Network) return diff --git a/src/model/service/task_service_test.go b/src/model/service/task_service_test.go new file mode 100644 index 00000000..82f087d3 --- /dev/null +++ b/src/model/service/task_service_test.go @@ -0,0 +1,92 @@ +package service + +import ( + "fmt" + "math/big" + "testing" + "time" + + "github.com/assimon/luuu/internal/testutil" + "github.com/assimon/luuu/model/dao" + "github.com/assimon/luuu/model/data" + "github.com/assimon/luuu/model/mdb" + "github.com/ethereum/go-ethereum/common" + "github.com/tidwall/gjson" +) + +func TestWalkTronGridPagesFollowsFingerprint(t *testing.T) { + var calls int + var ids []int64 + err := walkTronGridPages(func(query map[string]string) ([]byte, error) { + calls++ + switch calls { + case 1: + if query["fingerprint"] != "" { + t.Fatalf("unexpected first fingerprint: %q", query["fingerprint"]) + } + return []byte(`{"data":[{"id":1},{"id":2}],"meta":{"fingerprint":"next-page"}}`), nil + case 2: + if query["fingerprint"] != "next-page" { + t.Fatalf("unexpected second fingerprint: %q", query["fingerprint"]) + } + return []byte(`{"data":[{"id":3}],"meta":{}}`), nil + default: + return nil, fmt.Errorf("unexpected extra call %d", calls) + } + }, map[string]string{"limit": "100"}, func(record gjson.Result) bool { + ids = append(ids, record.Get("id").Int()) + return true + }) + if err != nil { + t.Fatalf("walkTronGridPages returned error: %v", err) + } + if calls != 2 { + t.Fatalf("expected 2 page fetches, got %d", calls) + } + if got := fmt.Sprint(ids); got != "[1 2 3]" { + t.Fatalf("unexpected collected ids %s", got) + } +} + +func TestTryProcessEvmERC20TransferSkipsHistoricalTransfer(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + order := &mdb.Orders{ + TradeId: "evm-history-order", + OrderId: "evm-history-order", + Amount: 1.00, + Currency: "USD", + ActualAmount: 1.00, + ReceiveAddress: "0x08c34c4e8b99e2503017ae09287bd0019b7096c6", + Token: "USDT", + Network: mdb.NetworkEthereum, + Status: mdb.StatusWaitPay, + } + if err := dao.Mdb.Create(order).Error; err != nil { + t.Fatalf("create evm order: %v", err) + } + if err := data.LockTransaction(order.Network, order.ReceiveAddress, order.Token, order.TradeId, order.ActualAmount, time.Hour); err != nil { + t.Fatalf("lock evm order: %v", err) + } + + TryProcessEvmERC20Transfer( + mdb.NetworkEthereum, + common.HexToAddress("0xdAC17F958D2ee523a2206206994597C13D831ec7"), + common.HexToAddress(order.ReceiveAddress), + big.NewInt(1_000_000), + "0xhistorical", + order.CreatedAt.TimestampMilli()-1000, + ) + + current, err := data.GetOrderInfoByTradeId(order.TradeId) + if err != nil { + t.Fatalf("reload evm order: %v", err) + } + if current.Status != mdb.StatusWaitPay { + t.Fatalf("historical transfer should not mark order paid: %+v", current) + } + if current.BlockTransactionId != "" { + t.Fatalf("historical transfer should not set block transaction id: %+v", current) + } +} diff --git a/src/mq/worker.go b/src/mq/worker.go index 3797c923..59339f0d 100644 --- a/src/mq/worker.go +++ b/src/mq/worker.go @@ -1,11 +1,11 @@ package mq import ( - "errors" "fmt" "io" "net/http" "net/url" + "strconv" "strings" "time" @@ -165,27 +165,51 @@ func processCallback(tradeID string) { } func sendOrderCallback(order *mdb.Orders) error { + checkCallbackResponse := func(statusCode int, body []byte, successBodies ...string) error { + if statusCode != http.StatusOK { + return fmt.Errorf("unexpected callback status: %d", statusCode) + } + normalized := strings.ToLower(strings.TrimSpace(string(body))) + for _, successBody := range successBodies { + if normalized == strings.ToLower(strings.TrimSpace(successBody)) { + return nil + } + } + return fmt.Errorf("unexpected callback body: %s", strings.TrimSpace(string(body))) + } switch order.PaymentType { case mdb.PaymentTypeEpay: - // 构造 EPay 标准回调参数 + paymentMerchantID := strings.TrimSpace(order.PaymentMerchantId) + if paymentMerchantID == "" && config.GetEpayPid() > 0 { + paymentMerchantID = strconv.Itoa(config.GetEpayPid()) + } + if paymentMerchantID == "" { + return fmt.Errorf("missing epay pid for trade_id=%s", order.TradeId) + } + pid, err := strconv.Atoi(paymentMerchantID) + if err != nil { + return err + } + paymentChannel := strings.TrimSpace(order.PaymentChannel) + if paymentChannel == "" { + paymentChannel = "alipay" + } notifyData := response.OrderNotifyResponseEpay{ - PID: config.GetEpayPid(), - TradeNo: order.TradeId, // epusdt 订单号作为 EPay 平台订单号 - OutTradeNo: order.OrderId, // 注意:EPay 回调要求商户订单号使用 out_trade_no 参数 - - Type: "alipay", + PID: pid, + TradeNo: order.TradeId, + OutTradeNo: order.OrderId, + Type: paymentChannel, Name: order.Name, Money: fmt.Sprintf("%.4f", order.Amount), TradeStatus: "TRADE_SUCCESS", } - signstr2, err := sign.Get(notifyData, config.GetEpayKey()) + signstr2, err := sign.Get(notifyData, config.GetEpaySignKey()) if err != nil { return err } - // 使用 form-encoded POST(EPay 标准协议格式) formData := url.Values{ "pid": {fmt.Sprintf("%d", notifyData.PID)}, "trade_no": {notifyData.TradeNo}, @@ -208,8 +232,9 @@ func sendOrderCallback(order *mdb.Orders) error { if err != nil { return err } - - fmt.Printf("notify_url response status: %d, body: %s\n", resp.StatusCode, string(responseBody)) + if err = checkCallbackResponse(resp.StatusCode, responseBody, "success", "ok"); err != nil { + return err + } default: @@ -221,6 +246,7 @@ func sendOrderCallback(order *mdb.Orders) error { ActualAmount: order.ActualAmount, ReceiveAddress: order.ReceiveAddress, Token: order.Token, + Network: order.Network, BlockTransactionId: order.BlockTransactionId, Status: mdb.StatusPaySuccess, } @@ -237,11 +263,8 @@ func sendOrderCallback(order *mdb.Orders) error { if err != nil { return err } - if resp.StatusCode() != http.StatusOK { - return errors.New(resp.Status()) - } - if string(resp.Body()) != "ok" { - return errors.New("not ok") + if err = checkCallbackResponse(resp.StatusCode(), resp.Body(), "ok"); err != nil { + return err } } diff --git a/src/mq/worker_test.go b/src/mq/worker_test.go index e3813c97..a2588def 100644 --- a/src/mq/worker_test.go +++ b/src/mq/worker_test.go @@ -1,6 +1,7 @@ package mq import ( + "encoding/json" "io" "net/http" "net/http/httptest" @@ -187,6 +188,75 @@ func TestDispatchPendingCallbacksHonorsBackoffAndPersistsSuccess(t *testing.T) { } } +func TestSendOrderCallbackUsesActualPaidTokenAndNetwork(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + var callbackBody map[string]interface{} + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := json.NewDecoder(r.Body).Decode(&callbackBody); err != nil { + t.Fatalf("decode callback body: %v", err) + } + _, _ = io.WriteString(w, "ok") + })) + defer server.Close() + + parent := &mdb.Orders{ + TradeId: "trade_parent_actual", + OrderId: "order_parent_actual", + Amount: 1.24, + Currency: "CNY", + ActualAmount: 0.18, + ReceiveAddress: "wallet_parent", + Token: "USDT", + Network: mdb.NetworkTron, + Status: mdb.StatusWaitPay, + NotifyUrl: server.URL, + CallBackConfirm: mdb.CallBackConfirmNo, + } + if err := dao.Mdb.Create(parent).Error; err != nil { + t.Fatalf("create parent order: %v", err) + } + sub := &mdb.Orders{ + TradeId: "trade_sub_actual", + OrderId: "order_sub_actual", + ParentTradeId: parent.TradeId, + Amount: 1.24, + Currency: "CNY", + ActualAmount: 0.17, + ReceiveAddress: "0x08c34c4e8b99e2503017ae09287bd0019b7096c6", + Token: "USDC", + Network: mdb.NetworkEthereum, + Status: mdb.StatusPaySuccess, + BlockTransactionId: "block_sub_actual", + CallBackConfirm: mdb.CallBackConfirmOk, + } + if err := dao.Mdb.Create(sub).Error; err != nil { + t.Fatalf("create sub-order: %v", err) + } + if _, err := data.MarkParentOrderSuccess(parent.TradeId, sub); err != nil { + t.Fatalf("mark parent success: %v", err) + } + + parent, err := data.GetOrderInfoByTradeId(parent.TradeId) + if err != nil { + t.Fatalf("reload parent order: %v", err) + } + if err = sendOrderCallback(parent); err != nil { + t.Fatalf("send callback: %v", err) + } + + if callbackBody["token"] != "USDC" { + t.Fatalf("expected callback token USDC, got %v", callbackBody["token"]) + } + if callbackBody["network"] != mdb.NetworkEthereum { + t.Fatalf("expected callback network %s, got %v", mdb.NetworkEthereum, callbackBody["network"]) + } + if callbackBody["receive_address"] != sub.ReceiveAddress { + t.Fatalf("expected callback receive address %s, got %v", sub.ReceiveAddress, callbackBody["receive_address"]) + } +} + func TestDispatchPendingCallbacksResumesRetryAfterRestart(t *testing.T) { cleanup := testutil.SetupTestDatabases(t) defer cleanup() @@ -259,6 +329,108 @@ func TestDispatchPendingCallbacksResumesRetryAfterRestart(t *testing.T) { } } +func TestDispatchPendingCallbacksEpayRequiresSuccessfulResponse(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + callbackLimiter = make(chan struct{}, 1) + callbackInflight = sync.Map{} + + var callbackType atomic.Value + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Fatalf("parse form: %v", err) + } + callbackType.Store(r.Form.Get("type")) + _, _ = io.WriteString(w, "fail") + })) + defer server.Close() + + order := &mdb.Orders{ + TradeId: "trade_callback_epay_fail", + OrderId: "order_callback_epay_fail", + Amount: 1, + Currency: "USD", + ActualAmount: 1, + ReceiveAddress: "wallet_epay_fail", + Token: "USDT", + Status: mdb.StatusPaySuccess, + NotifyUrl: server.URL, + CallbackNum: 0, + CallBackConfirm: mdb.CallBackConfirmNo, + PaymentType: mdb.PaymentTypeEpay, + PaymentChannel: "wxpay", + } + if err := dao.Mdb.Create(order).Error; err != nil { + t.Fatalf("create epay callback order: %v", err) + } + + dispatchPendingCallbacks() + + waitFor(t, 3*time.Second, func() bool { + current, err := data.GetOrderInfoByTradeId(order.TradeId) + if err != nil || current.ID <= 0 { + return false + } + return current.CallBackConfirm == mdb.CallBackConfirmNo && current.CallbackNum == 1 + }) + + if got, _ := callbackType.Load().(string); got != "wxpay" { + t.Fatalf("callback type = %q, want wxpay", got) + } +} + +func TestDispatchPendingCallbacksEpayAcceptsSuccessResponse(t *testing.T) { + cleanup := testutil.SetupTestDatabases(t) + defer cleanup() + + callbackLimiter = make(chan struct{}, 1) + callbackInflight = sync.Map{} + + var callbackType atomic.Value + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if err := r.ParseForm(); err != nil { + t.Fatalf("parse form: %v", err) + } + callbackType.Store(r.Form.Get("type")) + _, _ = io.WriteString(w, "success") + })) + defer server.Close() + + order := &mdb.Orders{ + TradeId: "trade_callback_epay_success", + OrderId: "order_callback_epay_success", + Amount: 1, + Currency: "USD", + ActualAmount: 1, + ReceiveAddress: "wallet_epay_success", + Token: "USDT", + Status: mdb.StatusPaySuccess, + NotifyUrl: server.URL, + CallbackNum: 0, + CallBackConfirm: mdb.CallBackConfirmNo, + PaymentType: mdb.PaymentTypeEpay, + PaymentChannel: "wxpay", + } + if err := dao.Mdb.Create(order).Error; err != nil { + t.Fatalf("create epay callback order: %v", err) + } + + dispatchPendingCallbacks() + + waitFor(t, 3*time.Second, func() bool { + current, err := data.GetOrderInfoByTradeId(order.TradeId) + if err != nil || current.ID <= 0 { + return false + } + return current.CallBackConfirm == mdb.CallBackConfirmOk && current.CallbackNum == 1 + }) + + if got, _ := callbackType.Load().(string); got != "wxpay" { + t.Fatalf("callback type = %q, want wxpay", got) + } +} + func waitFor(t *testing.T, timeout time.Duration, fn func() bool) { t.Helper() diff --git a/src/route/router.go b/src/route/router.go index d332acc1..2a110a13 100644 --- a/src/route/router.go +++ b/src/route/router.go @@ -7,6 +7,7 @@ import ( "io" "net/http" "strconv" + "strings" "github.com/assimon/luuu/config" "github.com/assimon/luuu/controller/comm" @@ -114,6 +115,19 @@ func RegisterRoute(e *echo.Echo) { return s } + pid := strings.TrimSpace(getString(params, "pid")) + if configuredPID := config.GetEpayPid(); configuredPID > 0 { + configured := strconv.Itoa(configuredPID) + if pid != "" && pid != configured { + return constant.SignatureErr + } + pid = configured + } + if pid == "" { + return constant.SignatureErr + } + params["pid"] = pid + signstr := getString(params, "sign") if signstr == "" { return constant.SignatureErr @@ -122,10 +136,7 @@ func RegisterRoute(e *echo.Echo) { delete(params, "sign") delete(params, "sign_type") - // we need to add pid to params for signature verification - params["pid"] = config.GetEpayPid() - - checkSignature, err := sign.Get(params, config.GetApiAuthToken()) + checkSignature, err := sign.Get(params, config.GetEpaySignKey()) if err != nil { return constant.SignatureErr } @@ -138,6 +149,19 @@ func RegisterRoute(e *echo.Echo) { notifyURL := getString(params, "notify_url") outTradeNo := getString(params, "out_trade_no") returnURL := getString(params, "return_url") + token := strings.TrimSpace(getString(params, "token")) + if token == "" { + token = config.GetEpayDefaultToken() + } + currency := strings.TrimSpace(getString(params, "currency")) + if currency == "" { + currency = config.GetEpayDefaultCurrency() + } + network := strings.TrimSpace(getString(params, "network")) + if network == "" { + network = config.GetEpayDefaultNetwork() + } + paymentChannel := strings.TrimSpace(getString(params, "type")) amountFloat, err := strconv.ParseFloat(money, 64) if err != nil { @@ -145,16 +169,18 @@ func RegisterRoute(e *echo.Echo) { } body := map[string]interface{}{ - "token": "usdt", - "currency": "cny", - "network": "tron", - "amount": amountFloat, - "notify_url": notifyURL, - "order_id": outTradeNo, - "redirect_url": returnURL, - "signature": signstr, - "name": name, - "payment_type": mdb.PaymentTypeEpay, + "token": token, + "currency": currency, + "network": network, + "amount": amountFloat, + "notify_url": notifyURL, + "order_id": outTradeNo, + "redirect_url": returnURL, + "signature": signstr, + "name": name, + "payment_type": mdb.PaymentTypeEpay, + "payment_channel": paymentChannel, + "payment_merchant_id": pid, } ctx.Set("request_body", body) diff --git a/src/route/router_test.go b/src/route/router_test.go index 350d098c..5a28975e 100644 --- a/src/route/router_test.go +++ b/src/route/router_test.go @@ -7,9 +7,11 @@ import ( "net/http/httptest" "net/url" "os" + "path/filepath" "strings" "testing" + "github.com/assimon/luuu/config" "github.com/assimon/luuu/model/dao" "github.com/assimon/luuu/model/mdb" "github.com/assimon/luuu/util/log" @@ -18,7 +20,16 @@ import ( "github.com/spf13/viper" ) -const testAPIToken = "test-secret-token" +const ( + testAPIToken = "test-secret-token" + testEpayKey = "test-epay-key" + testTronAddr1 = "TLa2f6VPqDgRE67v1736s7bJ8Ray5wYjU7" + testTronAddr2 = "TXLAQ63Xg1NAzckPwKHvzw7CSEmLMEqcdj" + testTronAddr3 = "TJRabPrwbZy45sbavfcjinPJC18kjpRTv8" + testSolAddr1 = "So11111111111111111111111111111111111111112" + testSolAddr2 = "11111111111111111111111111111111" + testEvmAddr1 = "0x08c34c4e8b99e2503017ae09287bd0019b7096c6" +) func setupTestEnv(t *testing.T) *echo.Echo { t.Helper() @@ -29,6 +40,7 @@ func setupTestEnv(t *testing.T) *echo.Echo { viper.Reset() viper.Set("db_type", "sqlite") viper.Set("api_auth_token", testAPIToken) + viper.Set("epay_key", testEpayKey) viper.Set("epay_pid", 1) viper.Set("app_uri", "http://localhost:8080") viper.Set("order_expiration_time", 10) @@ -40,6 +52,11 @@ func setupTestEnv(t *testing.T) *echo.Echo { viper.Set("runtime_sqlite_filename", tmpDir+"/runtime.db") log.Init() + wd, err := os.Getwd() + if err != nil { + t.Fatalf("Getwd: %v", err) + } + config.StaticFilePath = filepath.Join(wd, "..", "static") // init config paths os.Setenv("EPUSDT_CONFIG", tmpDir) @@ -57,8 +74,8 @@ func setupTestEnv(t *testing.T) *echo.Echo { dao.Mdb.AutoMigrate(&mdb.Orders{}, &mdb.WalletAddress{}, &mdb.SupportedAsset{}) // seed wallet addresses - dao.Mdb.Create(&mdb.WalletAddress{Network: mdb.NetworkTron, Address: "TTestTronAddress001", Status: mdb.TokenStatusEnable}) - dao.Mdb.Create(&mdb.WalletAddress{Network: mdb.NetworkSolana, Address: "SolTestAddress001", Status: mdb.TokenStatusEnable}) + dao.Mdb.Create(&mdb.WalletAddress{Network: mdb.NetworkTron, Address: testTronAddr1, Status: mdb.TokenStatusEnable}) + dao.Mdb.Create(&mdb.WalletAddress{Network: mdb.NetworkSolana, Address: testSolAddr1, Status: mdb.TokenStatusEnable}) // seed supported assets if empty var supportCnt int64 dao.Mdb.Model(&mdb.SupportedAsset{}).Count(&supportCnt) @@ -107,7 +124,7 @@ func signEpayValues(values url.Values) url.Values { } signParams[key] = items[0] } - sig, _ := sign.Get(signParams, testAPIToken) + sig, _ := sign.Get(signParams, testEpayKey) values.Set("sign", sig) values.Set("sign_type", "MD5") return values @@ -151,7 +168,7 @@ func TestCreateOrderEpusdtDefaultTron(t *testing.T) { if data["trade_id"] == nil || data["trade_id"] == "" { t.Error("expected trade_id in response") } - if data["receive_address"] != "TTestTronAddress001" { + if data["receive_address"] != testTronAddr1 { t.Errorf("expected tron address, got: %v", data["receive_address"]) } t.Logf("Order created: trade_id=%v address=%v amount=%v", data["trade_id"], data["receive_address"], data["actual_amount"]) @@ -190,7 +207,7 @@ func TestCreateOrderGmpayV1Solana(t *testing.T) { if data["trade_id"] == nil || data["trade_id"] == "" { t.Error("expected trade_id in response") } - if data["receive_address"] != "SolTestAddress001" { + if data["receive_address"] != testSolAddr1 { t.Errorf("expected solana address, got: %v", data["receive_address"]) } t.Logf("Order created: trade_id=%v address=%v amount=%v", data["trade_id"], data["receive_address"], data["actual_amount"]) @@ -366,7 +383,7 @@ func TestWalletAddAndList(t *testing.T) { // Add a solana wallet rec := doPostWithToken(e, "/payments/gmpay/v1/wallet/add", map[string]interface{}{ "network": "solana", - "address": "NewSolWallet001", + "address": testSolAddr2, }) t.Logf("Add: %s", rec.Body.String()) resp := parseResp(t, rec) @@ -377,7 +394,7 @@ func TestWalletAddAndList(t *testing.T) { // Add a tron wallet rec = doPostWithToken(e, "/payments/gmpay/v1/wallet/add", map[string]interface{}{ "network": "tron", - "address": "NewTronWallet001", + "address": testTronAddr2, }) resp = parseResp(t, rec) if resp["status_code"].(float64) != 200 { @@ -413,7 +430,7 @@ func TestWalletAddAndList(t *testing.T) { func TestWalletDuplicateRejected(t *testing.T) { e := setupTestEnv(t) - body := map[string]interface{}{"network": "solana", "address": "DupWallet001"} + body := map[string]interface{}{"network": "ethereum", "address": "0x08C34c4E8B99E2503017ae09287BD0019b7096C6"} rec := doPostWithToken(e, "/payments/gmpay/v1/wallet/add", body) resp := parseResp(t, rec) if resp["status_code"].(float64) != 200 { @@ -428,14 +445,29 @@ func TestWalletDuplicateRejected(t *testing.T) { } t.Logf("Duplicate rejected: %v", resp["message"]) - // Same address, different network should succeed rec = doPostWithToken(e, "/payments/gmpay/v1/wallet/add", map[string]interface{}{ - "network": "tron", - "address": "DupWallet001", + "network": "ethereum", + "address": testEvmAddr1, }) resp = parseResp(t, rec) - if resp["status_code"].(float64) != 200 { - t.Fatalf("same address on different network should succeed: %v", resp) + if resp["status_code"].(float64) == 200 { + t.Fatal("expected normalized EVM duplicate to be rejected") + } +} + +func TestWalletInvalidAddressRejected(t *testing.T) { + e := setupTestEnv(t) + + rec := doPostWithToken(e, "/payments/gmpay/v1/wallet/add", map[string]interface{}{ + "network": "tron", + "address": "invalid-tron-address", + }) + resp := parseResp(t, rec) + if resp["status_code"].(float64) == 200 { + t.Fatal("expected invalid wallet address to be rejected") + } + if resp["status_code"].(float64) != 10016 { + t.Fatalf("expected status_code=10016, got %v", resp["status_code"]) } } @@ -445,8 +477,8 @@ func TestWalletStatusAndDelete(t *testing.T) { // Add a wallet rec := doPostWithToken(e, "/payments/gmpay/v1/wallet/add", map[string]interface{}{ - "network": "solana", - "address": "StatusTestWallet", + "network": "ethereum", + "address": testEvmAddr1, }) resp := parseResp(t, rec) wallet := resp["data"].(map[string]interface{}) @@ -474,7 +506,7 @@ func TestWalletStatusAndDelete(t *testing.T) { wallets := resp["data"].([]interface{}) for _, w := range wallets { wm := w.(map[string]interface{}) - if wm["address"] == "StatusTestWallet" && wm["status"].(float64) != 2 { + if wm["address"] == testEvmAddr1 && wm["status"].(float64) != 2 { t.Error("wallet should be disabled") } } @@ -546,11 +578,65 @@ func TestCreateOrderNetworkIsolation(t *testing.T) { if !ok { t.Fatalf("expected data, got: %v", resp) } - if data["receive_address"] == "TTestTronAddress001" { + if data["receive_address"] == testTronAddr1 { t.Error("solana order should NOT get a tron address") } - if data["receive_address"] != "SolTestAddress001" { - t.Errorf("expected SolTestAddress001, got %v", data["receive_address"]) + if data["receive_address"] != testSolAddr1 { + t.Errorf("expected %s, got %v", testSolAddr1, data["receive_address"]) + } +} + +func TestCreateOrderRejectsDisabledSupportedAsset(t *testing.T) { + e := setupTestEnv(t) + + if err := dao.Mdb.Model(&mdb.SupportedAsset{}). + Where("network = ? AND token = ?", mdb.NetworkTron, "USDT"). + Update("status", mdb.TokenStatusDisable).Error; err != nil { + t.Fatalf("disable supported asset: %v", err) + } + + rec := doPost(e, "/payments/gmpay/v1/order/create-transaction", signBody(map[string]interface{}{ + "order_id": "unsupported-tron-usdt", + "amount": 1.00, + "token": "usdt", + "currency": "cny", + "network": "tron", + "notify_url": "http://localhost/notify", + })) + resp := parseResp(t, rec) + if resp["status_code"].(float64) != 10017 { + t.Fatalf("expected status_code=10017, got %v body=%v", resp["status_code"], resp) + } +} + +func TestSwitchNetworkRejectsDisabledSupportedAsset(t *testing.T) { + e := setupTestEnv(t) + + createResp := parseResp(t, doPost(e, "/payments/gmpay/v1/order/create-transaction", signBody(map[string]interface{}{ + "order_id": "switch-base-order", + "amount": 1.00, + "token": "usdt", + "currency": "cny", + "network": "tron", + "notify_url": "http://localhost/notify", + }))) + data := createResp["data"].(map[string]interface{}) + tradeID := data["trade_id"].(string) + + if err := dao.Mdb.Model(&mdb.SupportedAsset{}). + Where("network = ? AND token = ?", mdb.NetworkSolana, "USDT"). + Update("status", mdb.TokenStatusDisable).Error; err != nil { + t.Fatalf("disable supported asset: %v", err) + } + + rec := doPost(e, "/pay/switch-network", map[string]interface{}{ + "trade_id": tradeID, + "token": "usdt", + "network": "solana", + }) + resp := parseResp(t, rec) + if resp["status_code"].(float64) != 10017 { + t.Fatalf("expected status_code=10017, got %v body=%v", resp["status_code"], resp) } } @@ -579,6 +665,36 @@ func TestEpaySubmitPhpGetCompatible(t *testing.T) { } } +func TestCheckoutCounterInjectsPaymentOptions(t *testing.T) { + e := setupTestEnv(t) + + createResp := parseResp(t, doPost(e, "/payments/gmpay/v1/order/create-transaction", signBody(map[string]interface{}{ + "order_id": "checkout-options-order", + "amount": 1.00, + "token": "usdt", + "currency": "cny", + "network": "tron", + "notify_url": "http://localhost/notify", + }))) + data := createResp["data"].(map[string]interface{}) + tradeID := data["trade_id"].(string) + + rec := doGet(e, "/pay/checkout-counter/"+tradeID) + body := rec.Body.String() + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d body=%s", rec.Code, body) + } + if !strings.Contains(body, "var PAYMENT_OPTIONS = [") { + t.Fatalf("expected PAYMENT_OPTIONS injection, got body=%s", body) + } + if !strings.Contains(body, `"network":"tron"`) { + t.Fatalf("expected tron option in checkout html, got body=%s", body) + } + if !strings.Contains(body, `"network":"solana"`) { + t.Fatalf("expected solana option in checkout html, got body=%s", body) + } +} + func TestEpaySubmitPhpPostFormCompatible(t *testing.T) { e := setupTestEnv(t) @@ -602,3 +718,152 @@ func TestEpaySubmitPhpPostFormCompatible(t *testing.T) { t.Fatalf("expected checkout redirect, got %q", rec.Header().Get("Location")) } } + +func TestEpaySubmitPhpUsesEpayKeyAndStoresOverrides(t *testing.T) { + e := setupTestEnv(t) + + values := signEpayValues(url.Values{ + "pid": {"1"}, + "name": {"epay-override-001"}, + "type": {"wxpay"}, + "money": {"1.00"}, + "token": {"usdt"}, + "currency": {"usd"}, + "network": {"solana"}, + "out_trade_no": {"epay-override-001"}, + "notify_url": {"http://localhost/notify"}, + "return_url": {"http://localhost/return"}, + }) + + rec := doFormPost(e, "/payments/epay/v1/order/create-transaction/submit.php", values) + if rec.Code != http.StatusFound { + t.Fatalf("expected 302, got %d body=%s", rec.Code, rec.Body.String()) + } + + order := new(mdb.Orders) + if err := dao.Mdb.Where("order_id = ?", "epay-override-001").Limit(1).Find(order).Error; err != nil { + t.Fatalf("load created epay order: %v", err) + } + if order.ID == 0 { + t.Fatal("expected epay order to be created") + } + if order.Network != mdb.NetworkSolana || order.Token != "USDT" || order.Currency != "USD" { + t.Fatalf("unexpected stored payment mapping: %+v", order) + } + if order.PaymentType != mdb.PaymentTypeEpay { + t.Fatalf("unexpected payment type: %+v", order) + } + if order.PaymentChannel != "wxpay" { + t.Fatalf("unexpected payment channel: %+v", order) + } + if order.PaymentMerchantId != "1" { + t.Fatalf("unexpected payment merchant id: %+v", order) + } +} + +func TestSupportedAssetsFiltersByOrderContextAndValidWallets(t *testing.T) { + e := setupTestEnv(t) + + if err := dao.Mdb.Create(&mdb.WalletAddress{ + Network: mdb.NetworkBsc, + Address: "not-a-valid-evm-address", + Status: mdb.TokenStatusEnable, + }).Error; err != nil { + t.Fatalf("create invalid bsc wallet: %v", err) + } + + rec := doGet(e, "/payments/gmpay/v1/supported-assets?currency=cny&amount=1.00") + resp := parseResp(t, rec) + if resp["status_code"].(float64) != 200 { + t.Fatalf("supported assets request failed: %v", resp) + } + + data := resp["data"].(map[string]interface{}) + supports := data["supports"].([]interface{}) + body, _ := json.Marshal(supports) + bodyText := string(body) + if strings.Contains(bodyText, `"TRX"`) { + t.Fatalf("expected TRX to be filtered by rate availability, got %s", bodyText) + } + if strings.Contains(bodyText, `"bsc"`) { + t.Fatalf("expected bsc to be filtered because wallet is invalid, got %s", bodyText) + } + if !strings.Contains(bodyText, `"tron"`) || !strings.Contains(bodyText, `"USDT"`) { + t.Fatalf("expected tron/usdt to remain available, got %s", bodyText) + } +} + +func TestCreateTransactionAllowsMinimumAmount(t *testing.T) { + e := setupTestEnv(t) + + rec := doPost(e, "/payments/gmpay/v1/order/create-transaction", signBody(map[string]interface{}{ + "order_id": "minimum-amount-order", + "amount": 0.01, + "token": "usdt", + "currency": "usd", + "network": "tron", + "notify_url": "http://localhost/notify", + })) + resp := parseResp(t, rec) + if resp["status_code"].(float64) != 200 { + t.Fatalf("expected minimum amount to be accepted, got %v", resp) + } +} + +func TestCreateTransactionNormalizesStoredAmount(t *testing.T) { + e := setupTestEnv(t) + + rec := doPost(e, "/payments/gmpay/v1/order/create-transaction", signBody(map[string]interface{}{ + "order_id": "normalized-amount-order", + "amount": 1.239, + "token": "usdt", + "currency": "cny", + "network": "tron", + "notify_url": "http://localhost/notify", + })) + resp := parseResp(t, rec) + if resp["status_code"].(float64) != 200 { + t.Fatalf("create transaction failed: %v", resp) + } + data := resp["data"].(map[string]interface{}) + if data["amount"].(float64) != 1.24 { + t.Fatalf("expected normalized response amount 1.24, got %v", data["amount"]) + } + + order := new(mdb.Orders) + if err := dao.Mdb.Where("order_id = ?", "normalized-amount-order").Limit(1).Find(order).Error; err != nil { + t.Fatalf("load normalized order: %v", err) + } + if order.Amount != 1.24 { + t.Fatalf("expected normalized stored amount 1.24, got %+v", order) + } +} + +func TestCheckoutCounterInjectsTerminalOrderStatus(t *testing.T) { + e := setupTestEnv(t) + + createResp := parseResp(t, doPost(e, "/payments/gmpay/v1/order/create-transaction", signBody(map[string]interface{}{ + "order_id": "checkout-terminal-order", + "amount": 1.00, + "token": "usdt", + "currency": "cny", + "network": "tron", + "notify_url": "http://localhost/notify", + }))) + tradeID := createResp["data"].(map[string]interface{})["trade_id"].(string) + + if err := dao.Mdb.Model(&mdb.Orders{}). + Where("trade_id = ?", tradeID). + Update("status", mdb.StatusPaySuccess).Error; err != nil { + t.Fatalf("mark order paid: %v", err) + } + + rec := doGet(e, "/pay/checkout-counter/"+tradeID) + body := rec.Body.String() + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d body=%s", rec.Code, body) + } + if !strings.Contains(body, `status: "2"`) { + t.Fatalf("expected terminal order status injection, got body=%s", body) + } +} diff --git a/src/static/index.html b/src/static/index.html index 217d97ab..1cd4b432 100644 --- a/src/static/index.html +++ b/src/static/index.html @@ -531,16 +531,18 @@ amount: "{{.Amount}}", actualAmount: "{{.ActualAmount}}", token: "{{.Token}}", - currency: "{{.Currency}}", - network: "{{.Network}}", - receiveAddress: "{{.ReceiveAddress}}", - expirationTime: "{{.ExpirationTime}}", + currency: "{{.Currency}}", + network: "{{.Network}}", + status: "{{.Status}}", + receiveAddress: "{{.ReceiveAddress}}", + expirationTime: "{{.ExpirationTime}}", redirectUrl: "{{.RedirectUrl}}", createdAt: "{{.CreatedAt}}", is_selected: "{{.IsSelected}}" }; + var PAYMENT_OPTIONS = {{.PaymentOptionsJSON}}; - \ No newline at end of file + diff --git a/src/static/payment.js b/src/static/payment.js index 18ac7e50..4786e889 100644 --- a/src/static/payment.js +++ b/src/static/payment.js @@ -41,6 +41,8 @@ const API_ERRORS = { 10011: { en: 'Sub-order quantity exceeded', zh: '子订单数量超限', 'zh-hk': '子訂單數量超限', ja: 'サブ注文数が超過しました', ko: '하위 주문 수량 초과', ru: 'Превышено количество подзаказов' }, 10012: { en: 'Cannot switch network for sub-orders', zh: '不能对子订单切换网络', 'zh-hk': '不能對子訂單切換網絡', ja: 'サブ注文のネットワーク切替不可', ko: '하위 주문에 네트워크 전환 불가', ru: 'Нельзя переключить сеть для подзаказов' }, 10013: { en: 'Order is not in pending payment status', zh: '订单不是待支付状态', 'zh-hk': '訂單不是待支付狀態', ja: '注文は支払い待ち状態ではありません', ko: '주문이 결제 대기 상태가 아님', ru: 'Заказ не в статусе ожидания оплаты' }, + 10016: { en: 'Invalid wallet address', zh: '无效钱包地址', 'zh-hk': '無效錢包地址', ja: '無効なウォレットアドレス', ko: '유효하지 않은 지갑 주소', ru: 'Недопустимый адрес кошелька' }, + 10017: { en: 'Payment method unavailable', zh: '支付方式不可用', 'zh-hk': '支付方式不可用', ja: '支払い方法は利用できません', ko: '결제 수단을 사용할 수 없음', ru: 'Способ оплаты недоступен' }, }; /** @@ -77,9 +79,9 @@ async function apiFetch(url, opts = {}) { if (!fetchOpts.signal) fetchOpts.signal = AbortSignal.timeout(timeout); const res = await fetch(url, fetchOpts); const resp = await res.json().catch(() => null); - if (!res.ok || (resp?.code != null && resp.code !== 200)) { - const code = resp?.code ? Number(resp.code) : res.status; - throw new ApiError(code, getApiErrorMsg(code) ?? resp?.msg ?? `HTTP ${res.status}`); + const code = resp?.status_code != null ? Number(resp.status_code) : res.status; + if (!res.ok || (resp?.status_code != null && Number(resp.status_code) !== 200)) { + throw new ApiError(code, getApiErrorMsg(code) ?? resp?.message ?? `HTTP ${res.status}`); } return resp?.data ?? resp; } @@ -109,7 +111,13 @@ const CONFIG = { // ---- 后端接口 ---- api: { // 获取支持的网络和币种 - supportedAssets: () => '/payments/gmpay/v1/supported-assets', + supportedAssets: (currency, amount) => { + const qs = new URLSearchParams(); + if (currency) qs.set('currency', currency); + if (amount) qs.set('amount', amount); + const query = qs.toString(); + return query ? `/payments/gmpay/v1/supported-assets?${query}` : '/payments/gmpay/v1/supported-assets'; + }, // 切换网络接口:POST { trade_id, token, network },返回完整订单对象 selectMethod: () => '/pay/switch-network', // 轮询接口:GET,返回 { data: { status: number } } @@ -639,20 +647,46 @@ function setStepBar(step) { */ async function fetchSupportedAssets() { try { - const data = await apiFetch(CONFIG.api.supportedAssets()); + const data = await apiFetch(CONFIG.api.supportedAssets(ORDER.currency, ORDER.amount)); if (!data?.supports?.length) return null; return data.supports.flatMap(s => s.tokens.map(token => ({ token, network: s.network }))); } catch (e) { - console.warn('[supportedAssets]', e); return null; } } +function _normalizePaymentOptions(options) { + if (!Array.isArray(options)) return []; + const seen = new Set(); + return options.reduce((rows, option) => { + const token = String(option?.token ?? '').trim().toUpperCase(); + const network = String(option?.network ?? '').trim().toLowerCase(); + if (!token || !network) return rows; + const key = `${network}:${token}`; + if (seen.has(key)) return rows; + seen.add(key); + rows.push({ token, network }); + return rows; + }, []); +} + +async function ensurePaymentOptions() { + let options = _normalizePaymentOptions(window.PAYMENT_OPTIONS); + if (options.length) return options; + options = _normalizePaymentOptions(await fetchSupportedAssets()); + if (options.length) return options; + options = _normalizePaymentOptions([{ token: ORDER.token, network: ORDER.network }]); + return options; +} + let _step1Token = null; let _step1Opt = null; function initStep1() { - _step1Opt = PAYMENT_OPTIONS[0]; + _step1Opt = PAYMENT_OPTIONS.find(o => + o.network === String(ORDER.network || '').trim().toLowerCase() && + o.token === String(ORDER.token || '').trim().toUpperCase() + ) ?? PAYMENT_OPTIONS[0]; _step1Token = _step1Opt.token; _renderNetworkMenu(); _renderTokenMenu(); @@ -947,8 +981,12 @@ function _enterTerminalState(panel, { clearTimer = true, disableBtn = true } = { } function onPaymentSuccess() { + showSuccess(true); +} + +function showSuccess(redirect = false) { _enterTerminalState('success'); - if (ORDER?.redirectUrl && !ORDER.redirectUrl.startsWith('{{')) { + if (redirect && ORDER?.redirectUrl && !ORDER.redirectUrl.startsWith('{{')) { setTimeout(() => { window.location.href = ORDER.redirectUrl; }, CONFIG.redirect.delay); } } @@ -999,9 +1037,22 @@ document.addEventListener('DOMContentLoaded', async () => { setText('display-amount', `${ORDER.amount} ${ORDER.currency || ''}`); } - // 从 API 获取支持的网络和币种 - window.PAYMENT_OPTIONS = await fetchSupportedAssets(); - if (!PAYMENT_OPTIONS?.length) return; + const status = Number(ORDER?.status || 0); + const { paid, expired } = CONFIG.api.statusMap; + if (status === paid) { + showSuccess(false); + return; + } + if (status === expired) { + showExpired(); + return; + } + + window.PAYMENT_OPTIONS = await ensurePaymentOptions(); + if (!PAYMENT_OPTIONS?.length) { + showNotFound(); + return; + } // isselect=true 时跳过 Step 1,直接进入支付面板 const _isSelect = ORDER.is_selected && ORDER.is_selected !== 'false' && !ORDER.is_selected.startsWith('{{'); diff --git a/src/task/listen_bsc.go b/src/task/listen_bsc.go index 4990c7e2..4295c125 100644 --- a/src/task/listen_bsc.go +++ b/src/task/listen_bsc.go @@ -58,7 +58,7 @@ func StartBscWebSocketListener() { Topics: [][]common.Hash{}, } - runEvmWsLogListener("[BSC-WS]", wsURL, query, func(client *ethclient.Client, vLog types.Log) { + runEvmWsLogListener("[BSC-WS]", wsURL, query, evmBackfillLookbackBlocks(3*time.Second), func(client *ethclient.Client, vLog types.Log) { if len(vLog.Topics) < 3 { return } diff --git a/src/task/listen_eth.go b/src/task/listen_eth.go index d04badab..333b7dcf 100644 --- a/src/task/listen_eth.go +++ b/src/task/listen_eth.go @@ -62,7 +62,7 @@ func StartEthereumWebSocketListener() { Topics: [][]common.Hash{}, } - runEvmWsLogListener("[ETH-WS]", wsURL, query, func(client *ethclient.Client, vLog types.Log) { + runEvmWsLogListener("[ETH-WS]", wsURL, query, evmBackfillLookbackBlocks(12*time.Second), func(client *ethclient.Client, vLog types.Log) { if len(vLog.Topics) < 3 { return } diff --git a/src/task/listen_evm_ws.go b/src/task/listen_evm_ws.go index 087b1f71..7663a2d5 100644 --- a/src/task/listen_evm_ws.go +++ b/src/task/listen_evm_ws.go @@ -2,8 +2,10 @@ package task import ( "context" + "math/big" "time" + "github.com/assimon/luuu/config" "github.com/assimon/luuu/util/log" "github.com/ethereum/go-ethereum" @@ -11,13 +13,76 @@ import ( "github.com/ethereum/go-ethereum/ethclient" ) -func runEvmWsLogListener(logPrefix, wsURL string, query ethereum.FilterQuery, handleLog func(*ethclient.Client, types.Log)) { +func evmBackfillLookbackBlocks(avgBlockTime time.Duration) uint64 { + lookback := config.GetOrderExpirationTimeDuration() + 5*time.Minute + blocks := uint64(lookback / avgBlockTime) + if lookback%avgBlockTime != 0 { + blocks++ + } + blocks *= 2 + if blocks < 256 { + return 256 + } + return blocks +} + +func backfillEvmLogs(client *ethclient.Client, logPrefix string, query ethereum.FilterQuery, lastSeenBlock *uint64, lookbackBlocks uint64, handleLog func(*ethclient.Client, types.Log)) { + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + + header, err := client.HeaderByNumber(ctx, nil) + if err != nil { + log.Sugar.Warnf("%s latest header for backfill: %v", logPrefix, err) + return + } + latestBlock := header.Number.Uint64() + + startBlock := uint64(0) + switch { + case *lastSeenBlock > 0 && *lastSeenBlock < latestBlock: + startBlock = *lastSeenBlock + 1 + case latestBlock >= lookbackBlocks: + startBlock = latestBlock - lookbackBlocks + 1 + } + if startBlock > latestBlock { + return + } + + const chunkSize uint64 = 512 + for from := startBlock; from <= latestBlock; from += chunkSize { + to := from + chunkSize - 1 + if to > latestBlock { + to = latestBlock + } + filterQuery := query + filterQuery.FromBlock = big.NewInt(int64(from)) + filterQuery.ToBlock = big.NewInt(int64(to)) + + logs, err := client.FilterLogs(ctx, filterQuery) + if err != nil { + log.Sugar.Warnf("%s backfill logs %d-%d: %v", logPrefix, from, to, err) + return + } + for _, vLog := range logs { + if vLog.BlockNumber > *lastSeenBlock { + *lastSeenBlock = vLog.BlockNumber + } + handleLog(client, vLog) + } + } + if latestBlock > *lastSeenBlock { + *lastSeenBlock = latestBlock + } +} + +func runEvmWsLogListener(logPrefix, wsURL string, query ethereum.FilterQuery, lookbackBlocks uint64, handleLog func(*ethclient.Client, types.Log)) { const ( minBackoff = 2 * time.Second maxBackoff = 60 * time.Second rejoinWait = 3 * time.Second ) failWait := minBackoff + var lastSeenBlock uint64 for { client, err := ethclient.Dial(wsURL) @@ -41,13 +106,15 @@ func runEvmWsLogListener(logPrefix, wsURL string, query ethereum.FilterQuery, ha log.Sugar.Infof("%s connected, subscribed to USDT/USDC Transfer logs", logPrefix) - recvLoop(client, sub, logsCh, logPrefix, handleLog) + backfillEvmLogs(client, logPrefix, query, &lastSeenBlock, lookbackBlocks, handleLog) + + recvLoop(client, sub, logsCh, logPrefix, &lastSeenBlock, handleLog) time.Sleep(rejoinWait) } } -func recvLoop(client *ethclient.Client, sub ethereum.Subscription, logsCh <-chan types.Log, logPrefix string, handleLog func(*ethclient.Client, types.Log)) { +func recvLoop(client *ethclient.Client, sub ethereum.Subscription, logsCh <-chan types.Log, logPrefix string, lastSeenBlock *uint64, handleLog func(*ethclient.Client, types.Log)) { defer func() { sub.Unsubscribe() client.Close() @@ -67,6 +134,9 @@ func recvLoop(client *ethclient.Client, sub ethereum.Subscription, logsCh <-chan log.Sugar.Warnf("%s log channel closed, reconnecting", logPrefix) return } + if vLog.BlockNumber > *lastSeenBlock { + *lastSeenBlock = vLog.BlockNumber + } handleLog(client, vLog) } } diff --git a/src/task/listen_plasma.go b/src/task/listen_plasma.go index 9b447030..3adb3e77 100644 --- a/src/task/listen_plasma.go +++ b/src/task/listen_plasma.go @@ -54,7 +54,7 @@ func StartPlasmaWebSocketListener() { Topics: [][]common.Hash{}, } - runEvmWsLogListener("[PLASMA-WS]", wsURL, query, func(client *ethclient.Client, vLog types.Log) { + runEvmWsLogListener("[PLASMA-WS]", wsURL, query, evmBackfillLookbackBlocks(2*time.Second), func(client *ethclient.Client, vLog types.Log) { if len(vLog.Topics) < 3 { return } diff --git a/src/task/listen_polygon.go b/src/task/listen_polygon.go index 90aba847..d07dd729 100644 --- a/src/task/listen_polygon.go +++ b/src/task/listen_polygon.go @@ -60,7 +60,7 @@ func StartPolygonWebSocketListener() { Topics: [][]common.Hash{}, } - runEvmWsLogListener("[POLYGON-WS]", wsURL, query, func(client *ethclient.Client, vLog types.Log) { + runEvmWsLogListener("[POLYGON-WS]", wsURL, query, evmBackfillLookbackBlocks(2*time.Second), func(client *ethclient.Client, vLog types.Log) { if len(vLog.Topics) < 3 { return } diff --git a/src/telegram/handle.go b/src/telegram/handle.go index 95522df4..eb11d00a 100644 --- a/src/telegram/handle.go +++ b/src/telegram/handle.go @@ -9,6 +9,7 @@ import ( "github.com/assimon/luuu/model/data" "github.com/assimon/luuu/model/mdb" + "github.com/assimon/luuu/util/walletaddr" "github.com/gookit/goutil/mathutil" "github.com/gookit/goutil/strutil" tb "gopkg.in/telebot.v3" @@ -57,11 +58,11 @@ func OnTextMessageHandle(c tb.Context) error { } var err error - if !isValidAddressByNetwork(state.Network, msgText) { + if !walletaddr.Validate(state.Network, msgText) { _ = c.Send(fmt.Sprintf("钱包 [%s] 添加失败:不是合法的 %s 地址", msgText, strings.ToUpper(state.Network))) return nil } - storeAddress := normalizeWalletAddressByNetwork(state.Network, msgText) + storeAddress := walletaddr.Normalize(state.Network, msgText) _, err = data.AddWalletAddressWithNetwork(state.Network, storeAddress) if err != nil { return c.Send(err.Error()) diff --git a/src/telegram/utils.go b/src/telegram/utils.go deleted file mode 100644 index fde9c66f..00000000 --- a/src/telegram/utils.go +++ /dev/null @@ -1,80 +0,0 @@ -package telegram - -import ( - "crypto/sha256" - "encoding/hex" - "strings" - - "github.com/assimon/luuu/model/mdb" - "github.com/btcsuite/btcutil/base58" - "github.com/gagliardetto/solana-go" -) - -// isValidEthereumAddress 校验 0x + 20 字节十六进制(主网收款)。 -func isValidEthereumAddress(addr string) bool { - addr = strings.TrimSpace(addr) - if len(addr) != 42 || !strings.HasPrefix(addr, "0x") { - return false - } - _, err := hex.DecodeString(addr[2:]) - return err == nil -} - -// isValidTronAddress 校验 Tron Base58Check 地址是否合法 -func isValidTronAddress(addr string) bool { - // 基本过滤 - if len(addr) < 26 || len(addr) > 35 || addr[0] != 'T' { - return false - } - - decoded := base58.Decode(addr) - if len(decoded) != 25 { - return false - } - - // TRON 主网地址必须以 0x41 开头 - if decoded[0] != 0x41 { - return false - } - - // Base58Check 校验 - payload := decoded[:21] // 前 21 字节 - checksum := decoded[21:] // 后 4 字节 - - hash := sha256.Sum256(payload) - hash2 := sha256.Sum256(hash[:]) - - return string(checksum) == string(hash2[:4]) -} - -func isValidAddressByNetwork(network, addr string) bool { - switch strings.ToLower(strings.TrimSpace(network)) { - case mdb.NetworkTron: - return isValidTronAddress(addr) - case mdb.NetworkSolana: - return isValidSolanaAddress(addr) - default: - // 其余 EVM 链统一使用 0x 地址校验 - return isValidEthereumAddress(addr) - } -} - -func normalizeWalletAddressByNetwork(network, addr string) string { - addr = strings.TrimSpace(addr) - switch strings.ToLower(strings.TrimSpace(network)) { - case mdb.NetworkTron, mdb.NetworkSolana: - return addr - default: - return strings.ToLower(addr) - } -} - -// isValidSolanaAddress 校验 Solana Base58 地址是否合法(32 字节公钥)。 -func isValidSolanaAddress(addr string) bool { - addr = strings.TrimSpace(addr) - if len(addr) < 32 || len(addr) > 44 { - return false - } - _, err := solana.PublicKeyFromBase58(addr) - return err == nil -} diff --git a/src/util/constant/errno.go b/src/util/constant/errno.go index d38ca17b..a692bb5b 100644 --- a/src/util/constant/errno.go +++ b/src/util/constant/errno.go @@ -18,26 +18,30 @@ var Errno = map[int]string{ 10013: "order is not awaiting payment", 10014: "supported asset already exists", 10015: "supported asset not found", + 10016: "invalid wallet address", + 10017: "payment method unavailable", } var ( - SystemErr = Err(400) - SignatureErr = Err(401) - WalletAddressAlreadyExists = Err(10001) - OrderAlreadyExists = Err(10002) - NotAvailableWalletAddress = Err(10003) - PayAmountErr = Err(10004) - NotAvailableAmountErr = Err(10005) - RateAmountErr = Err(10006) - OrderBlockAlreadyProcess = Err(10007) - OrderNotExists = Err(10008) - ParamsMarshalErr = Err(10009) - OrderStatusConflict = Err(10010) - SubOrderLimitExceeded = Err(10011) - CannotSwitchSubOrder = Err(10012) - OrderNotWaitPay = Err(10013) + SystemErr = Err(400) + SignatureErr = Err(401) + WalletAddressAlreadyExists = Err(10001) + OrderAlreadyExists = Err(10002) + NotAvailableWalletAddress = Err(10003) + PayAmountErr = Err(10004) + NotAvailableAmountErr = Err(10005) + RateAmountErr = Err(10006) + OrderBlockAlreadyProcess = Err(10007) + OrderNotExists = Err(10008) + ParamsMarshalErr = Err(10009) + OrderStatusConflict = Err(10010) + SubOrderLimitExceeded = Err(10011) + CannotSwitchSubOrder = Err(10012) + OrderNotWaitPay = Err(10013) SupportedAssetAlreadyExists = Err(10014) SupportedAssetNotFound = Err(10015) + InvalidWalletAddress = Err(10016) + PaymentMethodUnavailable = Err(10017) ) type RspError struct { diff --git a/src/util/walletaddr/walletaddr.go b/src/util/walletaddr/walletaddr.go new file mode 100644 index 00000000..ec978aca --- /dev/null +++ b/src/util/walletaddr/walletaddr.go @@ -0,0 +1,79 @@ +package walletaddr + +import ( + "crypto/sha256" + "encoding/hex" + "strings" + + "github.com/assimon/luuu/model/mdb" + "github.com/btcsuite/btcutil/base58" + "github.com/gagliardetto/solana-go" +) + +func NormalizeNetwork(network string) string { + return strings.ToLower(strings.TrimSpace(network)) +} + +func IsEVMNetwork(network string) bool { + switch NormalizeNetwork(network) { + case mdb.NetworkEthereum, mdb.NetworkBsc, mdb.NetworkPolygon, mdb.NetworkPlasma: + return true + } + return false +} + +func Normalize(network, address string) string { + address = strings.TrimSpace(address) + if IsEVMNetwork(network) { + return strings.ToLower(address) + } + return address +} + +func Validate(network, address string) bool { + address = strings.TrimSpace(address) + switch NormalizeNetwork(network) { + case mdb.NetworkTron: + return isValidTronAddress(address) + case mdb.NetworkSolana: + return isValidSolanaAddress(address) + case mdb.NetworkEthereum, mdb.NetworkBsc, mdb.NetworkPolygon, mdb.NetworkPlasma: + return isValidEVMAddress(address) + default: + return false + } +} + +func isValidEVMAddress(address string) bool { + if len(address) != 42 || !strings.HasPrefix(address, "0x") { + return false + } + _, err := hex.DecodeString(address[2:]) + return err == nil +} + +func isValidTronAddress(address string) bool { + if len(address) < 26 || len(address) > 35 || address[0] != 'T' { + return false + } + decoded := base58.Decode(address) + if len(decoded) != 25 { + return false + } + if decoded[0] != 0x41 { + return false + } + payload := decoded[:21] + checksum := decoded[21:] + hash := sha256.Sum256(payload) + hash2 := sha256.Sum256(hash[:]) + return string(checksum) == string(hash2[:4]) +} + +func isValidSolanaAddress(address string) bool { + if len(address) < 32 || len(address) > 44 { + return false + } + _, err := solana.PublicKeyFromBase58(address) + return err == nil +} diff --git a/wiki/API.md b/wiki/API.md index abefe1b4..3b31239d 100644 --- a/wiki/API.md +++ b/wiki/API.md @@ -71,7 +71,7 @@ signature : 1cd4b52df5587cfb1968b0c0c6e156cd ## POST 创建交易 -POST /api/v1/order/create-transaction +POST /payments/gmpay/v1/order/create-transaction > Body 请求参数 @@ -91,7 +91,7 @@ POST /api/v1/order/create-transaction |---|---|---|---|-----------|---------------| |body|body|object| 否 || | |» order_id|body|string| 是 | 请求支付订单号 | | -|» amount|body|number| 是 | 支付金额(CNY) | 小数点保留后2位,最少0.01 | +|» amount|body|number| 是 | 支付金额 | 小数点保留后2位,且最少需能兑换 0.01 USDT;若使用 USD,则最少 0.01 | |» notify_url|body|string| 是 | 异步回调地址 | | |» redirect_url|body|string| 否 | 同步跳转地址 || |» signature|body|string| 是 | 签名 | 接口统一加密方式 |