diff --git a/controller/relay.go b/controller/relay.go index 65fe6fbe0cf2..8b282f3c3a70 100644 --- a/controller/relay.go +++ b/controller/relay.go @@ -150,21 +150,25 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { relayInfo.SetEstimatePromptTokens(tokens) - priceData, err := helper.ModelPriceHelper(c, relayInfo, tokens, meta) - if err != nil { - newAPIError = types.NewError(err, types.ErrorCodeModelPriceError, types.ErrOptionWithStatusCode(http.StatusBadRequest)) - return - } - - // common.SetContextKey(c, constant.ContextKeyTokenCountMeta, meta) - - if priceData.FreeModel { - logger.LogInfo(c, fmt.Sprintf("模型 %s 免费,跳过预扣费", relayInfo.OriginModelName)) + if relaycommon.IsClaudeCountTokensRequest(relayInfo) { + relayInfo.FinalPreConsumedQuota = 0 } else { - newAPIError = service.PreConsumeBilling(c, priceData.QuotaToPreConsume, relayInfo) - if newAPIError != nil { + priceData, err := helper.ModelPriceHelper(c, relayInfo, tokens, meta) + if err != nil { + newAPIError = types.NewError(err, types.ErrorCodeModelPriceError, types.ErrOptionWithStatusCode(http.StatusBadRequest)) return } + + // common.SetContextKey(c, constant.ContextKeyTokenCountMeta, meta) + + if priceData.FreeModel { + logger.LogInfo(c, fmt.Sprintf("模型 %s 免费,跳过预扣费", relayInfo.OriginModelName)) + } else { + newAPIError = service.PreConsumeBilling(c, priceData.QuotaToPreConsume, relayInfo) + if newAPIError != nil { + return + } + } } defer func() { diff --git a/controller/subscription_payment_epay.go b/controller/subscription_payment_epay.go index 7dece6badce2..303353b91d5a 100644 --- a/controller/subscription_payment_epay.go +++ b/controller/subscription_payment_epay.go @@ -84,10 +84,20 @@ func SubscriptionRequestEpay(c *gin.Context) { return } + group, err := model.GetUserGroup(userId, true) + if err != nil { + c.JSON(http.StatusOK, gin.H{"message": "error", "data": "获取用户分组失败"}) + return + } + payMoney := getPayMoney(int64(plan.PriceAmount), group) + if payMoney < 0.01 { + common.ApiErrorMsg(c, "套餐金额过低") + return + } order := &model.SubscriptionOrder{ UserId: userId, PlanId: plan.Id, - Money: plan.PriceAmount, + Money: payMoney, TradeNo: tradeNo, PaymentMethod: req.PaymentMethod, PaymentProvider: model.PaymentProviderEpay, @@ -102,7 +112,7 @@ func SubscriptionRequestEpay(c *gin.Context) { Type: req.PaymentMethod, ServiceTradeNo: tradeNo, Name: fmt.Sprintf("SUB:%s", plan.Title), - Money: strconv.FormatFloat(plan.PriceAmount, 'f', 2, 64), + Money: strconv.FormatFloat(payMoney, 'f', 2, 64), Device: epay.PC, NotifyUrl: notifyUrl, ReturnUrl: returnUrl, diff --git a/controller/subscription_payment_waffo_pancake.go b/controller/subscription_payment_waffo_pancake.go index 0915ddc6b240..7370e4d4c776 100644 --- a/controller/subscription_payment_waffo_pancake.go +++ b/controller/subscription_payment_waffo_pancake.go @@ -79,10 +79,15 @@ func SubscriptionRequestWaffoPancakePay(c *gin.Context) { // dispatch in WaffoPancakeWebhook. tradeNo := fmt.Sprintf("WAFFO_PANCAKE_SUB-%d-%d-%s", userId, time.Now().UnixMilli(), randstr.String(6)) + payMoney := getWaffoPancakePayMoney(int64(plan.PriceAmount), user.Group) + if payMoney < 0.01 { + c.JSON(http.StatusOK, gin.H{"message": "error", "data": "套餐金额过低"}) + return + } order := &model.SubscriptionOrder{ UserId: userId, PlanId: plan.Id, - Money: plan.PriceAmount, + Money: payMoney, TradeNo: tradeNo, PaymentMethod: model.PaymentMethodWaffoPancake, PaymentProvider: model.PaymentProviderWaffoPancake, @@ -100,7 +105,7 @@ func SubscriptionRequestWaffoPancakePay(c *gin.Context) { ProductID: plan.WaffoPancakeProductId, BuyerIdentity: service.WaffoPancakeBuyerIdentityFromUserID(user.Id), PriceSnapshot: &service.WaffoPancakePriceSnapshot{ - Amount: decimal.NewFromFloat(plan.PriceAmount).StringFixed(2), + Amount: decimal.NewFromFloat(payMoney).StringFixed(2), TaxCategory: "saas", }, BuyerEmail: getWaffoPancakeBuyerEmail(user), @@ -114,7 +119,7 @@ func SubscriptionRequestWaffoPancakePay(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"message": "error", "data": "拉起支付失败"}) return } - logger.LogInfo(c.Request.Context(), fmt.Sprintf("Waffo Pancake 订阅订单创建成功 user_id=%d plan_id=%d trade_no=%s session_id=%s money=%.2f", userId, plan.Id, tradeNo, session.SessionID, plan.PriceAmount)) + logger.LogInfo(c.Request.Context(), fmt.Sprintf("Waffo Pancake 订阅订单创建成功 user_id=%d plan_id=%d trade_no=%s session_id=%s money=%.2f", userId, plan.Id, tradeNo, session.SessionID, payMoney)) c.JSON(http.StatusOK, gin.H{ "message": "success", diff --git a/relay/channel/claude/adaptor.go b/relay/channel/claude/adaptor.go index 6daf5b6f245e..c7a3e927c84b 100644 --- a/relay/channel/claude/adaptor.go +++ b/relay/channel/claude/adaptor.go @@ -42,7 +42,11 @@ func (a *Adaptor) Init(info *relaycommon.RelayInfo) { } func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { - requestURL := fmt.Sprintf("%s/v1/messages", info.ChannelBaseUrl) + requestPath := "/v1/messages" + if relaycommon.IsClaudeCountTokensRequest(info) { + requestPath = "/v1/messages/count_tokens" + } + requestURL := fmt.Sprintf("%s%s", info.ChannelBaseUrl, requestPath) if !shouldAppendClaudeBetaQuery(info) { return requestURL, nil } diff --git a/relay/channel/claude/adaptor_test.go b/relay/channel/claude/adaptor_test.go new file mode 100644 index 000000000000..2510c283f543 --- /dev/null +++ b/relay/channel/claude/adaptor_test.go @@ -0,0 +1,33 @@ +package claude + +import ( + "testing" + + relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/stretchr/testify/require" +) + +func TestGetRequestURLUsesMessagesCountTokensPath(t *testing.T) { + adaptor := &Adaptor{} + info := &relaycommon.RelayInfo{ + ChannelBaseUrl: "https://api.anthropic.com", + RelayMode: relayconstant.RelayModeClaudeCountTokens, + } + + requestURL, err := adaptor.GetRequestURL(info) + require.NoError(t, err) + require.Equal(t, "https://api.anthropic.com/v1/messages/count_tokens", requestURL) +} + +func TestGetRequestURLKeepsMessagesPathForClaudeMessages(t *testing.T) { + adaptor := &Adaptor{} + info := &relaycommon.RelayInfo{ + ChannelBaseUrl: "https://api.anthropic.com", + RequestURLPath: "/v1/messages", + } + + requestURL, err := adaptor.GetRequestURL(info) + require.NoError(t, err) + require.Equal(t, "https://api.anthropic.com/v1/messages", requestURL) +} diff --git a/relay/claude_handler.go b/relay/claude_handler.go index 527363205a1f..96c6d63bc799 100644 --- a/relay/claude_handler.go +++ b/relay/claude_handler.go @@ -132,6 +132,12 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ } } + if relaycommon.IsClaudeCountTokensRequest(info) { + request.MaxTokens = nil + request.MaxTokensToSample = nil + request.Stream = nil + } + if !model_setting.GetGlobalSettings().PassThroughRequestEnabled && !info.ChannelSetting.PassThroughBodyEnabled && service.ShouldChatCompletionsUseResponsesGlobal(info.ChannelId, info.ChannelType, info.OriginModelName) { @@ -211,6 +217,16 @@ func ClaudeHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *typ } } + if relaycommon.IsClaudeCountTokensRequest(info) { + respBody, err := io.ReadAll(httpResp.Body) + service.CloseResponseBodyGracefully(httpResp) + if err != nil { + return types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + service.IOCopyBytesGracefully(c, httpResp, respBody) + return nil + } + usage, newAPIError := adaptor.DoResponse(c, httpResp, info) if newAPIError != nil { // reset status code 重置状态码 diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index fa52e05674a2..5af6f74376ef 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -4,6 +4,7 @@ import ( "encoding/json" "errors" "fmt" + "net/url" "strconv" "strings" "time" @@ -535,6 +536,24 @@ func cloneRequestHeaders(c *gin.Context) map[string]string { return headers } +func IsClaudeCountTokensRequestPath(path string) bool { + parsedPath := path + if parsedURL, err := url.Parse(path); err == nil && parsedURL.Path != "" { + parsedPath = parsedURL.Path + } + return parsedPath == "/v1/messages/count_tokens" +} + +func IsClaudeCountTokensRequest(info *RelayInfo) bool { + if info == nil { + return false + } + if info.RelayMode == relayconstant.RelayModeClaudeCountTokens { + return true + } + return IsClaudeCountTokensRequestPath(info.RequestURLPath) +} + func GenRelayInfo(c *gin.Context, relayFormat types.RelayFormat, request dto.Request, ws *websocket.Conn) (*RelayInfo, error) { var info *RelayInfo var err error diff --git a/relay/common/relay_info_test.go b/relay/common/relay_info_test.go index e53ec804ca06..8df30a3317c6 100644 --- a/relay/common/relay_info_test.go +++ b/relay/common/relay_info_test.go @@ -3,6 +3,7 @@ package common import ( "testing" + relayconstant "github.com/QuantumNous/new-api/relay/constant" "github.com/QuantumNous/new-api/types" "github.com/stretchr/testify/require" ) @@ -38,3 +39,10 @@ func TestRelayInfoGetFinalRequestRelayFormatNilReceiver(t *testing.T) { var info *RelayInfo require.Equal(t, types.RelayFormat(""), info.GetFinalRequestRelayFormat()) } + +func TestIsClaudeCountTokensRequest(t *testing.T) { + require.True(t, IsClaudeCountTokensRequest(&RelayInfo{RequestURLPath: "/v1/messages/count_tokens"})) + require.True(t, IsClaudeCountTokensRequest(&RelayInfo{RequestURLPath: "/v1/messages/count_tokens?beta=true"})) + require.True(t, IsClaudeCountTokensRequest(&RelayInfo{RelayMode: relayconstant.RelayModeClaudeCountTokens})) + require.False(t, IsClaudeCountTokensRequest(&RelayInfo{RequestURLPath: "/v1/messages"})) +} diff --git a/relay/constant/relay_mode.go b/relay/constant/relay_mode.go index 256715679213..fa821a99a94e 100644 --- a/relay/constant/relay_mode.go +++ b/relay/constant/relay_mode.go @@ -52,6 +52,8 @@ const ( RelayModeGemini RelayModeResponsesCompact + + RelayModeClaudeCountTokens ) func Path2RelayMode(path string) int { @@ -86,6 +88,8 @@ func Path2RelayMode(path string) int { relayMode = RelayModeRerank } else if strings.HasPrefix(path, "/v1/realtime") { relayMode = RelayModeRealtime + } else if strings.HasPrefix(path, "/v1/messages/count_tokens") { + relayMode = RelayModeClaudeCountTokens } else if strings.HasPrefix(path, "/v1beta/models") || strings.HasPrefix(path, "/v1/models") { relayMode = RelayModeGemini } else if strings.HasPrefix(path, "/mj") { diff --git a/relay/constant/relay_mode_test.go b/relay/constant/relay_mode_test.go new file mode 100644 index 000000000000..279d1f13ff64 --- /dev/null +++ b/relay/constant/relay_mode_test.go @@ -0,0 +1,11 @@ +package constant + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestPath2RelayModeClaudeCountTokens(t *testing.T) { + require.Equal(t, RelayModeClaudeCountTokens, Path2RelayMode("/v1/messages/count_tokens")) +} diff --git a/router/relay-router.go b/router/relay-router.go index 17a13cad7fd6..b404339ae841 100644 --- a/router/relay-router.go +++ b/router/relay-router.go @@ -85,6 +85,9 @@ func SetRelayRouter(router *gin.Engine) { httpRouter.Use(middleware.Distribute()) // claude related routes + httpRouter.POST("/messages/count_tokens", func(c *gin.Context) { + controller.Relay(c, types.RelayFormatClaude) + }) httpRouter.POST("/messages", func(c *gin.Context) { controller.Relay(c, types.RelayFormatClaude) })