diff --git a/common/api_type.go b/common/api_type.go index c198ffc0..94cf7391 100644 --- a/common/api_type.go +++ b/common/api_type.go @@ -77,6 +77,8 @@ func ChannelType2APIType(channelType int) (int, bool) { apiType = constant.APITypeCodex case constant.ChannelTypeAdvancedCustom: apiType = constant.APITypeAdvancedCustom + case constant.ChannelTypeSub2API: + apiType = constant.APITypeSub2API } if apiType == -1 { return constant.APITypeOpenAI, false diff --git a/common/endpoint_defaults.go b/common/endpoint_defaults.go index 11ec7921..80ed7c80 100644 --- a/common/endpoint_defaults.go +++ b/common/endpoint_defaults.go @@ -20,6 +20,7 @@ var defaultEndpointInfoMap = map[constant.EndpointType]EndpointInfo{ constant.EndpointTypeOpenAI: {Path: "/v1/chat/completions", Method: "POST"}, constant.EndpointTypeOpenAIResponse: {Path: "/v1/responses", Method: "POST"}, constant.EndpointTypeOpenAIResponseCompact: {Path: "/v1/responses/compact", Method: "POST"}, + constant.EndpointTypeOpenAIAlphaSearch: {Path: "/v1/alpha/search", Method: "POST"}, constant.EndpointTypeAnthropic: {Path: "/v1/messages", Method: "POST"}, constant.EndpointTypeGemini: {Path: "/v1beta/models/{model}:generateContent", Method: "POST"}, constant.EndpointTypeJinaRerank: {Path: "/v1/rerank", Method: "POST"}, diff --git a/common/endpoint_type.go b/common/endpoint_type.go index a5e2ff84..fe3eefcd 100644 --- a/common/endpoint_type.go +++ b/common/endpoint_type.go @@ -30,6 +30,20 @@ func GetEndpointTypesByChannelType(channelType int, modelName string) []constant endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAI, constant.EndpointTypeOpenAIResponse} case constant.ChannelTypeSora: endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAIVideo} + case constant.ChannelTypeSub2API: + endpointTypes = []constant.EndpointType{ + constant.EndpointTypeOpenAI, + constant.EndpointTypeOpenAIResponse, + constant.EndpointTypeAnthropic, + constant.EndpointTypeGemini, + constant.EndpointTypeOpenAIAlphaSearch, + } + case constant.ChannelTypeCodex: + endpointTypes = []constant.EndpointType{ + constant.EndpointTypeOpenAIResponse, + constant.EndpointTypeOpenAIResponseCompact, + constant.EndpointTypeOpenAIAlphaSearch, + } default: if IsOpenAIResponseOnlyModel(modelName) { endpointTypes = []constant.EndpointType{constant.EndpointTypeOpenAIResponse} diff --git a/constant/api_type.go b/constant/api_type.go index f3657a11..a65d915b 100644 --- a/constant/api_type.go +++ b/constant/api_type.go @@ -37,5 +37,6 @@ const ( APITypeReplicate APITypeCodex APITypeAdvancedCustom + APITypeSub2API APITypeDummy // this one is only for count, do not add any channel after this ) diff --git a/constant/channel.go b/constant/channel.go index 45ec9a44..d9903f4f 100644 --- a/constant/channel.go +++ b/constant/channel.go @@ -56,6 +56,7 @@ const ( ChannelTypeReplicate = 56 ChannelTypeCodex = 57 ChannelTypeAdvancedCustom = 58 + ChannelTypeSub2API = 59 ChannelTypeDummy // this one is only for count, do not add any channel after this ) @@ -120,6 +121,7 @@ var ChannelBaseURLs = []string{ "https://api.replicate.com", //56 "https://chatgpt.com", //57 "", //58 + "", //59 } var ChannelTypeNames = map[int]string{ @@ -178,6 +180,7 @@ var ChannelTypeNames = map[int]string{ ChannelTypeReplicate: "Replicate", ChannelTypeCodex: "ChatGPT Subscription (Codex)", ChannelTypeAdvancedCustom: "Advanced Custom", + ChannelTypeSub2API: "Sub2API", } func GetChannelTypeName(channelType int) string { diff --git a/constant/endpoint_type.go b/constant/endpoint_type.go index 8681bf06..5d27956d 100644 --- a/constant/endpoint_type.go +++ b/constant/endpoint_type.go @@ -6,6 +6,7 @@ const ( EndpointTypeOpenAI EndpointType = "openai" EndpointTypeOpenAIResponse EndpointType = "openai-response" EndpointTypeOpenAIResponseCompact EndpointType = "openai-response-compact" + EndpointTypeOpenAIAlphaSearch EndpointType = "openai-alpha-search" EndpointTypeAnthropic EndpointType = "anthropic" EndpointTypeGemini EndpointType = "gemini" EndpointTypeJinaRerank EndpointType = "jina-rerank" diff --git a/controller/option.go b/controller/option.go index 6a91086a..0367c793 100644 --- a/controller/option.go +++ b/controller/option.go @@ -235,6 +235,15 @@ func UpdateOption(c *gin.Context) { }) return } + case operation_setting.ToolPriceOptionKey: + err = operation_setting.ValidateToolPricesJSON(option.Value.(string)) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "success": false, + "message": err.Error(), + }) + return + } case "ImageRatio": err = ratio_setting.UpdateImageRatioByJSONString(option.Value.(string)) if err != nil { diff --git a/controller/relay.go b/controller/relay.go index 6e91ccb6..a214ea3f 100644 --- a/controller/relay.go +++ b/controller/relay.go @@ -49,6 +49,8 @@ func relayHandler(c *gin.Context, info *relaycommon.RelayInfo) *types.NewAPIErro err = relay.EmbeddingHelper(c, info) case relayconstant.RelayModeResponses, relayconstant.RelayModeResponsesCompact: err = relay.ResponsesHelper(c, info) + case relayconstant.RelayModeAlphaSearch: + err = relay.AlphaSearchHelper(c, info) default: err = relay.TextHelper(c, info) } diff --git a/dto/alpha_search_request.go b/dto/alpha_search_request.go new file mode 100644 index 00000000..7628a6a6 --- /dev/null +++ b/dto/alpha_search_request.go @@ -0,0 +1,39 @@ +package dto + +import ( + "encoding/json" + + "github.com/QuantumNous/new-api/types" + + "github.com/gin-gonic/gin" +) + +// AlphaSearchRequest is the Codex standalone web search request. +// RawBody preserves the original JSON so unknown fields are forwarded intact. +type AlphaSearchRequest struct { + Model string `json:"model"` + Id string `json:"id,omitempty"` + Stream *bool `json:"stream,omitempty"` + RawBody json.RawMessage `json:"-"` +} + +func (r *AlphaSearchRequest) GetTokenCountMeta() *types.TokenCountMeta { + combineText := "" + if len(r.RawBody) > 0 { + combineText = string(r.RawBody) + } + return &types.TokenCountMeta{ + CombineText: combineText, + TokenType: types.TokenTypeTokenizer, + } +} + +func (r *AlphaSearchRequest) IsStream(c *gin.Context) bool { + return false +} + +func (r *AlphaSearchRequest) SetModelName(modelName string) { + if modelName != "" { + r.Model = modelName + } +} diff --git a/dto/channel_settings.go b/dto/channel_settings.go index c92a3f98..b938521c 100644 --- a/dto/channel_settings.go +++ b/dto/channel_settings.go @@ -106,6 +106,7 @@ const ( advancedCustomEndpointPathOpenAIChat = "/v1/chat/completions" advancedCustomEndpointPathOpenAIResponses = "/v1/responses" advancedCustomEndpointPathOpenAIResponsesCompact = "/v1/responses/compact" + advancedCustomEndpointPathOpenAIAlphaSearch = "/v1/alpha/search" advancedCustomEndpointPathClaudeMessages = "/v1/messages" advancedCustomEndpointPathJinaRerank = "/v1/rerank" advancedCustomEndpointPathImageGeneration = "/v1/images/generations" @@ -204,6 +205,8 @@ func advancedCustomEndpointTypeFromIncomingPath(incomingPath string) (constant.E return constant.EndpointTypeOpenAIResponse, true case advancedCustomEndpointPathOpenAIResponsesCompact: return constant.EndpointTypeOpenAIResponseCompact, true + case advancedCustomEndpointPathOpenAIAlphaSearch: + return constant.EndpointTypeOpenAIAlphaSearch, true case advancedCustomEndpointPathClaudeMessages: return constant.EndpointTypeAnthropic, true case advancedCustomEndpointPathJinaRerank: @@ -475,6 +478,12 @@ func validateAdvancedCustomUpstreamTarget(index int, upstreamPath string) error } func validateAdvancedCustomConverterPath(index int, incomingPath string, converter string) error { + if incomingPath == advancedCustomEndpointPathOpenAIAlphaSearch { + if converter == advancedCustomConverterNone { + return nil + } + return fmt.Errorf("advanced_custom.advanced_routes[%d].converter does not match incoming_path: %s", index, converter) + } switch converter { case advancedCustomConverterNone: return nil diff --git a/dto/channel_settings_test.go b/dto/channel_settings_test.go index 080863be..9b42516e 100644 --- a/dto/channel_settings_test.go +++ b/dto/channel_settings_test.go @@ -477,3 +477,45 @@ func TestAdvancedCustomSupportedEndpointTypesForModel(t *testing.T) { constant.EndpointTypeAnthropic, }, config.SupportedEndpointTypesForModel("other-model")) } + +func TestAdvancedCustomValidateAlphaSearchConverterPath(t *testing.T) { + valid := &AdvancedCustomConfig{ + Routes: []AdvancedCustomRoute{ + { + IncomingPath: "/v1/alpha/search", + UpstreamPath: "/v1/alpha/search", + Converter: advancedCustomConverterNone, + }, + }, + } + require.NoError(t, valid.Validate()) + assert.Equal(t, []constant.EndpointType{ + constant.EndpointTypeOpenAIAlphaSearch, + }, valid.SupportedEndpointTypesForModel("gpt-5.1")) + + nonNoneConverters := []string{ + advancedCustomConverterClaudeMessagesToOpenAIChat, + advancedCustomConverterOpenAIChatToClaudeMessages, + advancedCustomConverterOpenAIChatToOpenAIResponses, + advancedCustomConverterOpenAIResponsesToOpenAIChat, + advancedCustomConverterOpenAIResponsesToGemini, + advancedCustomConverterGeminiContentToOpenAIChat, + advancedCustomConverterOpenAIChatToGeminiContent, + } + for _, converter := range nonNoneConverters { + t.Run(converter, func(t *testing.T) { + config := &AdvancedCustomConfig{ + Routes: []AdvancedCustomRoute{ + { + IncomingPath: "/v1/alpha/search", + UpstreamPath: "/v1/alpha/search", + Converter: converter, + }, + }, + } + err := config.Validate() + require.Error(t, err) + assert.Contains(t, err.Error(), "converter does not match incoming_path") + }) + } +} diff --git a/dto/gemini.go b/dto/gemini.go index e9bbb7f2..ab285dfa 100644 --- a/dto/gemini.go +++ b/dto/gemini.go @@ -438,10 +438,15 @@ func (c *GeminiChatGenerationConfig) UnmarshalJSON(data []byte) error { type MediaResolution string type GeminiChatCandidate struct { - Content GeminiChatContent `json:"content"` - FinishReason *string `json:"finishReason"` - Index int64 `json:"index"` - SafetyRatings []GeminiChatSafetyRating `json:"safetyRatings"` + Content GeminiChatContent `json:"content"` + FinishReason *string `json:"finishReason"` + Index int64 `json:"index"` + SafetyRatings []GeminiChatSafetyRating `json:"safetyRatings"` + GroundingMetadata *GeminiGroundingMetadata `json:"groundingMetadata,omitempty"` +} + +type GeminiGroundingMetadata struct { + WebSearchQueries []string `json:"webSearchQueries,omitempty"` } type GeminiChatSafetyRating struct { diff --git a/dto/openai_response.go b/dto/openai_response.go index 2de6014f..7959f074 100644 --- a/dto/openai_response.go +++ b/dto/openai_response.go @@ -320,42 +320,6 @@ func (o *OpenAIResponsesResponse) GetOpenAIError() *types.OpenAIError { return GetOpenAIError(o.Error) } -func (o *OpenAIResponsesResponse) HasImageGenerationCall() bool { - if len(o.Output) == 0 { - return false - } - for _, output := range o.Output { - if output.Type == ResponsesOutputTypeImageGenerationCall { - return true - } - } - return false -} - -func (o *OpenAIResponsesResponse) GetQuality() string { - if len(o.Output) == 0 { - return "" - } - for _, output := range o.Output { - if output.Type == ResponsesOutputTypeImageGenerationCall { - return output.Quality - } - } - return "" -} - -func (o *OpenAIResponsesResponse) GetSize() string { - if len(o.Output) == 0 { - return "" - } - for _, output := range o.Output { - if output.Type == ResponsesOutputTypeImageGenerationCall { - return output.Size - } - } - return "" -} - type IncompleteDetails struct { Reason string `json:"reason"` } @@ -368,6 +332,7 @@ type ResponsesOutput struct { Content []ResponsesOutputContent `json:"content"` Quality string `json:"quality"` Size string `json:"size"` + Result string `json:"result,omitempty"` CallId string `json:"call_id,omitempty"` Name string `json:"name,omitempty"` Arguments json.RawMessage `json:"arguments,omitempty"` @@ -399,11 +364,17 @@ type ResponsesReasoningSummaryPart struct { const ( BuildInToolWebSearchPreview = "web_search_preview" + BuildInToolWebSearch = "web_search" BuildInToolFileSearch = "file_search" + BuildInToolGoogleSearch = "google_search" + BuildInToolImageGeneration = "image_generation" ) const ( - BuildInCallWebSearchCall = "web_search_call" + BuildInCallWebSearchCall = "web_search_call" + BuildInCallFileSearchCall = "file_search_call" + BuildInCallFunctionCall = "function_call" + BuildInCallToolUse = "tool_use" ) const ( diff --git a/model/option.go b/model/option.go index 89a233ec..7ab64e0d 100644 --- a/model/option.go +++ b/model/option.go @@ -204,7 +204,17 @@ func SyncOptions(frequency int) { } } +func validateOptionValue(key string, value string) error { + if key == operation_setting.ToolPriceOptionKey { + return operation_setting.ValidateToolPricesJSON(value) + } + return nil +} + func UpdateOption(key string, value string) error { + if err := validateOptionValue(key, value); err != nil { + return err + } // Save to database first option := Option{ Key: key, @@ -229,6 +239,11 @@ func UpdateOptionsBulk(values map[string]string) error { if len(values) == 0 { return nil } + for key, value := range values { + if err := validateOptionValue(key, value); err != nil { + return err + } + } err := DB.Transaction(func(tx *gorm.DB) error { for k, v := range values { option := Option{Key: k} @@ -584,6 +599,11 @@ func updateOptionMap(key string, value string) (err error) { // handleConfigUpdate 处理分层配置更新,返回是否已处理 func handleConfigUpdate(key, value string) bool { + if key == operation_setting.ToolPriceOptionKey { + operation_setting.LoadToolPricesFromJSONString(value) + return true + } + parts := strings.SplitN(key, ".", 2) if len(parts) != 2 { return false // 不是分层配置 @@ -607,8 +627,6 @@ func handleConfigUpdate(key, value string) bool { // 特定配置的后处理 if configName == "performance_setting" { performance_setting.UpdateAndSync() - } else if configName == "tool_price_setting" { - operation_setting.RebuildToolPriceIndex() } else if configName == "billing_setting" { InvalidatePricingCache() ratio_setting.InvalidateExposedDataCache() diff --git a/relay/alpha_search_handler.go b/relay/alpha_search_handler.go new file mode 100644 index 00000000..299408fa --- /dev/null +++ b/relay/alpha_search_handler.go @@ -0,0 +1,136 @@ +package relay + +import ( + "errors" + "fmt" + "io" + "net/http" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" + "github.com/QuantumNous/new-api/logger" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/relay/helper" + "github.com/QuantumNous/new-api/service" + "github.com/QuantumNous/new-api/types" + + "github.com/gin-gonic/gin" +) + +func AlphaSearchHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { + info.InitChannelMeta(c) + + switch info.ChannelType { + case constant.ChannelTypeSub2API, constant.ChannelTypeCodex, constant.ChannelTypeAdvancedCustom: + default: + // Allow retry onto another channel that may support this endpoint. + return types.NewError( + errors.New("channel does not support /v1/alpha/search"), + types.ErrorCodeInvalidRequest, + ) + } + + request, ok := info.Request.(*dto.AlphaSearchRequest) + if !ok { + return types.NewErrorWithStatusCode( + fmt.Errorf("invalid request type, expected *dto.AlphaSearchRequest, got %T", info.Request), + types.ErrorCodeInvalidRequest, + http.StatusBadRequest, + types.ErrOptionWithSkipRetry(), + ) + } + + err := helper.ModelMappedHelper(c, info, request) + if err != nil { + return types.NewError(err, types.ErrorCodeChannelModelMappedError, types.ErrOptionWithSkipRetry()) + } + + jsonData, err := buildAlphaSearchRequestBody(request.RawBody, info.OriginModelName, info.UpstreamModelName) + if err != nil { + return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) + } + + if len(info.ParamOverride) > 0 { + jsonData, err = relaycommon.ApplyParamOverrideWithRelayInfo(jsonData, info) + if err != nil { + return newAPIErrorFromParamOverride(err) + } + } + + logger.LogDebug(c, "requestBody: %s", jsonData) + body, size, closer, err := relaycommon.NewOutboundJSONBody(jsonData) + if err != nil { + return types.NewError(err, types.ErrorCodeConvertRequestFailed, types.ErrOptionWithSkipRetry()) + } + defer closer.Close() + info.UpstreamRequestBodySize = size + + adaptor := GetAdaptor(info.ApiType) + if adaptor == nil { + return types.NewError(fmt.Errorf("invalid api type: %d", info.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry()) + } + adaptor.Init(info) + + resp, err := adaptor.DoRequest(c, info, body) + if err != nil { + return types.NewOpenAIError(err, types.ErrorCodeDoRequestFailed, http.StatusInternalServerError) + } + + statusCodeMappingStr := c.GetString("status_code_mapping") + httpResp, ok := resp.(*http.Response) + if !ok || httpResp == nil { + return types.NewOpenAIError(errors.New("invalid http response"), types.ErrorCodeDoRequestFailed, http.StatusInternalServerError) + } + defer httpResp.Body.Close() + + if httpResp.StatusCode < 200 || httpResp.StatusCode >= 300 { + newAPIError = service.RelayErrorHandler(c.Request.Context(), httpResp, false) + service.ResetStatusCode(newAPIError, statusCodeMappingStr) + return newAPIError + } + + if contentType := httpResp.Header.Get("Content-Type"); contentType != "" { + c.Writer.Header().Set("Content-Type", contentType) + } + c.Writer.WriteHeader(httpResp.StatusCode) + if _, err := io.Copy(c.Writer, httpResp.Body); err != nil { + return types.NewError(err, types.ErrorCodeDoRequestFailed, types.ErrOptionWithSkipRetry()) + } + + // Upstream alpha search returns no usage; bill one web_search_preview call. + if info.ResponsesUsageInfo == nil { + info.ResponsesUsageInfo = &relaycommon.ResponsesUsageInfo{ + BuiltInTools: make(map[string]*relaycommon.BuildInToolInfo), + } + } + if info.ResponsesUsageInfo.BuiltInTools == nil { + info.ResponsesUsageInfo.BuiltInTools = make(map[string]*relaycommon.BuildInToolInfo) + } + info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview] = &relaycommon.BuildInToolInfo{ + ToolName: dto.BuildInToolWebSearchPreview, + CallCount: 1, + } + + usage := &dto.Usage{} + service.PostTextConsumeQuota(c, info, usage, nil) + return nil +} + +// buildAlphaSearchRequestBody returns RawBody unchanged unless the model was +// mapped, in which case only the "model" field is rewritten so unknown fields +// are preserved. +func buildAlphaSearchRequestBody(rawBody []byte, originModel, upstreamModel string) ([]byte, error) { + if len(rawBody) == 0 { + return nil, errors.New("empty alpha search request body") + } + if upstreamModel == "" || upstreamModel == originModel { + return rawBody, nil + } + var body map[string]any + if err := common.Unmarshal(rawBody, &body); err != nil { + return nil, err + } + body["model"] = upstreamModel + return common.Marshal(body) +} diff --git a/relay/alpha_search_handler_test.go b/relay/alpha_search_handler_test.go new file mode 100644 index 00000000..5847bad0 --- /dev/null +++ b/relay/alpha_search_handler_test.go @@ -0,0 +1,47 @@ +package relay + +import ( + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestBuildAlphaSearchRequestBodyPreservesUnknownFields(t *testing.T) { + raw := []byte(`{ + "id":"req_1", + "model":"gpt-5.1", + "input":[{"role":"user","content":"hi"}], + "commands":{"search_query":[{"q":"weather","recency":1}]}, + "settings":{"locale":"en"}, + "future_field":{"nested":true} + }`) + + out, err := buildAlphaSearchRequestBody(raw, "gpt-5.1", "gpt-5.1-mapped") + require.NoError(t, err) + + var body map[string]any + require.NoError(t, common.Unmarshal(out, &body)) + assert.Equal(t, "gpt-5.1-mapped", body["model"]) + assert.Equal(t, "req_1", body["id"]) + require.Contains(t, body, "commands") + require.Contains(t, body, "settings") + require.Contains(t, body, "future_field") + require.Contains(t, body, "input") + + commands, ok := body["commands"].(map[string]any) + require.True(t, ok) + require.Contains(t, commands, "search_query") + + future, ok := body["future_field"].(map[string]any) + require.True(t, ok) + assert.Equal(t, true, future["nested"]) +} + +func TestBuildAlphaSearchRequestBodyNoMappingKeepsRawBytes(t *testing.T) { + raw := []byte(`{"model":"gpt-5.1","commands":{"search_query":[{"q":"x"}]},"future_field":1}`) + out, err := buildAlphaSearchRequestBody(raw, "gpt-5.1", "gpt-5.1") + require.NoError(t, err) + assert.Equal(t, raw, out) +} diff --git a/relay/channel/claude/relay-claude.go b/relay/channel/claude/relay-claude.go index 8488cc9b..8f9c74b3 100644 --- a/relay/channel/claude/relay-claude.go +++ b/relay/channel/claude/relay-claude.go @@ -114,6 +114,7 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud data = patchClaudeMessageDeltaUsageData(data, buildMessageDeltaPatchUsage(&claudeResponse, claudeInfo)) } } + countClaudeStreamBillableTools(c, info, &claudeResponse) helper.ClaudeChunkData(c, claudeResponse, data) } else if info.RelayFormat == types.RelayFormatOpenAI { response := StreamResponseClaude2OpenAI(&claudeResponse) @@ -122,6 +123,8 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud return nil } + countClaudeStreamBillableTools(c, info, &claudeResponse) + err = helper.ObjectData(c, response) if err != nil { logger.LogError(c, "send_stream_response_failed: "+err.Error()) @@ -130,6 +133,23 @@ func HandleStreamResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud return nil } +func countClaudeStreamBillableTools(c *gin.Context, info *relaycommon.RelayInfo, claudeResponse *dto.ClaudeResponse) { + if claudeResponse == nil { + return + } + if claudeResponse.Type == "content_block_start" && + claudeResponse.ContentBlock != nil && + claudeResponse.ContentBlock.Type == "tool_use" { + info.CountBillableToolCall(dto.BuildInCallToolUse, claudeResponse.ContentBlock.Name) + } + if claudeResponse.Type == "message_delta" && + claudeResponse.Usage != nil && + claudeResponse.Usage.ServerToolUse != nil && + claudeResponse.Usage.ServerToolUse.WebSearchRequests > 0 { + c.Set("claude_web_search_requests", claudeResponse.Usage.ServerToolUse.WebSearchRequests) + } +} + func HandleStreamFinalResponse(c *gin.Context, info *relaycommon.RelayInfo, claudeInfo *ClaudeResponseInfo) { if claudeInfo.Usage.PromptTokens == 0 { //上游出错 @@ -235,6 +255,12 @@ func HandleClaudeResponseData(c *gin.Context, info *relaycommon.RelayInfo, claud c.Set("claude_web_search_requests", claudeResponse.Usage.ServerToolUse.WebSearchRequests) } + for _, block := range claudeResponse.Content { + if block.Type == "tool_use" { + info.CountBillableToolCall(dto.BuildInCallToolUse, block.Name) + } + } + service.IOCopyBytesGracefully(c, httpResp, responseData) return nil } diff --git a/relay/channel/claude/tool_billing_test.go b/relay/channel/claude/tool_billing_test.go new file mode 100644 index 00000000..3a780c7d --- /dev/null +++ b/relay/channel/claude/tool_billing_test.go @@ -0,0 +1,76 @@ +package claude + +import ( + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/dto" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/setting/operation_setting" + "github.com/QuantumNous/new-api/types" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestHandleClaudeResponseDataCountsToolUse(t *testing.T) { + gin.SetMode(gin.TestMode) + operation_setting.SetToolPriceForTest("lookup_fn", 3.0) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest("lookup_fn") + }) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + info := &relaycommon.RelayInfo{ + OriginModelName: "claude-3-7-sonnet", + RelayFormat: types.RelayFormatClaude, + } + claudeInfo := &ClaudeResponseInfo{Usage: &dto.Usage{}} + + data := []byte(`{ + "type":"message", + "content":[ + {"type":"text","text":"hi"}, + {"type":"tool_use","id":"tu1","name":"lookup_fn","input":{}}, + {"type":"server_tool_use","id":"stu1","name":"web_search","input":{}} + ], + "usage":{"input_tokens":1,"output_tokens":1} + }`) + + err := HandleClaudeResponseData(c, info, claudeInfo, nil, data) + require.Nil(t, err) + require.NotNil(t, info.ResponsesUsageInfo) + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, "lookup_fn") + assert.Equal(t, 1, info.ResponsesUsageInfo.BuiltInTools["lookup_fn"].CallCount) + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, "web_search") +} + +func TestCountClaudeStreamBillableToolsSetsWebSearchRequests(t *testing.T) { + gin.SetMode(gin.TestMode) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + info := &relaycommon.RelayInfo{OriginModelName: "claude-3-7-sonnet"} + + countClaudeStreamBillableTools(c, info, &dto.ClaudeResponse{ + Type: "message_delta", + Usage: &dto.ClaudeUsage{ + ServerToolUse: &dto.ClaudeServerToolUse{WebSearchRequests: 3}, + }, + }) + assert.Equal(t, 3, c.GetInt("claude_web_search_requests")) + + operation_setting.SetToolPriceForTest("stream_fn", 2.0) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest("stream_fn") + }) + countClaudeStreamBillableTools(c, info, &dto.ClaudeResponse{ + Type: "content_block_start", + ContentBlock: &dto.ClaudeMediaMessage{ + Type: "tool_use", + Name: "stream_fn", + }, + }) + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, "stream_fn") + assert.Equal(t, 1, info.ResponsesUsageInfo.BuiltInTools["stream_fn"].CallCount) +} diff --git a/relay/channel/codex/adaptor.go b/relay/channel/codex/adaptor.go index ef4d4fa0..45ec08d5 100644 --- a/relay/channel/codex/adaptor.go +++ b/relay/channel/codex/adaptor.go @@ -112,18 +112,20 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request } func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) { - if info.RelayMode != relayconstant.RelayModeResponses && info.RelayMode != relayconstant.RelayModeResponsesCompact { + switch info.RelayMode { + case relayconstant.RelayModeAlphaSearch: + // Alpha search responses are handled by relay.AlphaSearchHelper. + return nil, types.NewError(errors.New("codex channel: alpha search response should be handled by AlphaSearchHelper"), types.ErrorCodeInvalidRequest) + case relayconstant.RelayModeResponsesCompact: + return openai.OaiResponsesCompactionHandler(c, resp) + case relayconstant.RelayModeResponses: + if info.IsStream { + return openai.OaiResponsesStreamHandler(c, info, resp) + } + return openai.OaiResponsesHandler(c, info, resp) + default: return nil, types.NewError(errors.New("codex channel: endpoint not supported"), types.ErrorCodeInvalidRequest) } - - if info.RelayMode == relayconstant.RelayModeResponsesCompact { - return openai.OaiResponsesCompactionHandler(c, resp) - } - - if info.IsStream { - return openai.OaiResponsesStreamHandler(c, info, resp) - } - return openai.OaiResponsesHandler(c, info, resp) } func (a *Adaptor) GetModelList() []string { @@ -135,12 +137,16 @@ func (a *Adaptor) GetChannelName() string { } func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { - if info.RelayMode != relayconstant.RelayModeResponses && info.RelayMode != relayconstant.RelayModeResponsesCompact { - return "", errors.New("codex channel: only /v1/responses and /v1/responses/compact are supported") - } - path := "/backend-api/codex/responses" - if info.RelayMode == relayconstant.RelayModeResponsesCompact { + var path string + switch info.RelayMode { + case relayconstant.RelayModeResponses: + path = "/backend-api/codex/responses" + case relayconstant.RelayModeResponsesCompact: path = "/backend-api/codex/responses/compact" + case relayconstant.RelayModeAlphaSearch: + path = "/backend-api/codex/alpha/search" + default: + return "", errors.New("codex channel: only /v1/responses, /v1/responses/compact and /v1/alpha/search are supported") } return relaycommon.GetFullRequestURL(info.ChannelBaseUrl, path, info.ChannelType), nil } diff --git a/relay/channel/codex/adaptor_test.go b/relay/channel/codex/adaptor_test.go new file mode 100644 index 00000000..8bbb5dff --- /dev/null +++ b/relay/channel/codex/adaptor_test.go @@ -0,0 +1,26 @@ +package codex + +import ( + "testing" + + "github.com/QuantumNous/new-api/constant" + relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGetRequestURLAlphaSearch(t *testing.T) { + adaptor := &Adaptor{} + info := &relaycommon.RelayInfo{ + ChannelMeta: &relaycommon.ChannelMeta{ + ChannelType: constant.ChannelTypeCodex, + ChannelBaseUrl: "https://chatgpt.com", + }, + RelayMode: relayconstant.RelayModeAlphaSearch, + } + + url, err := adaptor.GetRequestURL(info) + require.NoError(t, err) + assert.Equal(t, "https://chatgpt.com/backend-api/codex/alpha/search", url) +} diff --git a/relay/channel/gemini/relay-gemini.go b/relay/channel/gemini/relay-gemini.go index 556d7dc2..5052c4cc 100644 --- a/relay/channel/gemini/relay-gemini.go +++ b/relay/channel/gemini/relay-gemini.go @@ -76,6 +76,18 @@ func geminiResponseUsageText(response *dto.GeminiChatResponse) string { return text.String() } +func markGeminiGoogleSearchCall(c *gin.Context, response *dto.GeminiChatResponse) { + if c == nil || response == nil { + return + } + for _, candidate := range response.Candidates { + if candidate.GroundingMetadata != nil && len(candidate.GroundingMetadata.WebSearchQueries) > 0 { + c.Set("gemini_google_search_call", true) + return + } + } +} + func buildUsageFromGeminiResponse(c *gin.Context, info *relaycommon.RelayInfo, response *dto.GeminiChatResponse) dto.Usage { metadata := response.GetUsageMetadata() if dto.HasGeminiUsageMetadataTokens(metadata) { @@ -149,6 +161,8 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason)) } + markGeminiGoogleSearchCall(c, &geminiResponse) + // 统计图片数量 for _, candidate := range geminiResponse.Candidates { for _, part := range candidate.Content.Parts { @@ -308,6 +322,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R if err != nil { return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } + markGeminiGoogleSearchCall(c, &geminiResponse) if len(geminiResponse.Candidates) == 0 { usage := buildUsageFromGeminiResponse(c, info, &geminiResponse) diff --git a/relay/channel/gemini/relay_responses.go b/relay/channel/gemini/relay_responses.go index 08f96ab4..5c987962 100644 --- a/relay/channel/gemini/relay_responses.go +++ b/relay/channel/gemini/relay_responses.go @@ -31,6 +31,7 @@ func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h if err := common.Unmarshal(responseBody, &geminiResponse); err != nil { return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) } + markGeminiGoogleSearchCall(c, &geminiResponse) if len(geminiResponse.Candidates) == 0 { usage := buildUsageFromGeminiResponse(c, info, &geminiResponse) if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil { diff --git a/relay/channel/openai/relay-openai.go b/relay/channel/openai/relay-openai.go index 50415c8b..17cad566 100644 --- a/relay/channel/openai/relay-openai.go +++ b/relay/channel/openai/relay-openai.go @@ -119,6 +119,8 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re var usage = &dto.Usage{} var lastStreamData string var secondLastStreamData string // 存储倒数第二个stream data,用于音频模型 + seenStreamToolCalls := make(map[string]struct{}) + var streamFunctionCallNames []string // 检查是否为音频模型 isAudioModel := strings.Contains(strings.ToLower(model), "audio") @@ -137,6 +139,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re } lastStreamData = data + collectStreamFunctionCallNames(data, seenStreamToolCalls, &streamFunctionCallNames) if err := processTokenData(info.RelayMode, data, &responseTextBuilder, &toolCount); err != nil { logger.LogError(c, "error processing stream token data: "+err.Error()) sr.Error(err) @@ -182,11 +185,40 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re applyUsagePostProcessing(info, usage, common.StringToByteSlice(lastStreamData)) + for _, name := range streamFunctionCallNames { + info.CountBillableToolCall(dto.BuildInCallFunctionCall, name) + } + HandleFinalResponse(c, info, lastStreamData, responseId, createAt, model, systemFingerprint, usage, containStreamUsage) return usage, nil } +func collectStreamFunctionCallNames(data string, seen map[string]struct{}, names *[]string) { + var streamResponse dto.ChatCompletionsStreamResponse + if err := common.UnmarshalJsonStr(data, &streamResponse); err != nil { + return + } + for _, choice := range streamResponse.Choices { + for i, tc := range choice.Delta.ToolCalls { + name := tc.Function.Name + if name == "" { + continue + } + toolIdx := i + if tc.Index != nil { + toolIdx = *tc.Index + } + key := fmt.Sprintf("%d-%d", choice.Index, toolIdx) + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + *names = append(*names, name) + } + } +} + func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { defer service.CloseResponseBodyGracefully(resp) @@ -228,6 +260,12 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo } } + for _, choice := range simpleResponse.Choices { + for _, tc := range choice.Message.ParseToolCalls() { + info.CountBillableToolCall(dto.BuildInCallFunctionCall, tc.Function.Name) + } + } + forceFormat := false if info.ChannelSetting.ForceFormat { forceFormat = true diff --git a/relay/channel/openai/relay_responses.go b/relay/channel/openai/relay_responses.go index 92931831..697cc965 100644 --- a/relay/channel/openai/relay_responses.go +++ b/relay/channel/openai/relay_responses.go @@ -34,12 +34,6 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http return nil, types.WithOpenAIError(*oaiError, resp.StatusCode) } - if responsesResponse.HasImageGenerationCall() { - c.Set("image_generation_call", true) - c.Set("image_generation_call_quality", responsesResponse.GetQuality()) - c.Set("image_generation_call_size", responsesResponse.GetSize()) - } - // 写入新的 response body service.IOCopyBytesGracefully(c, resp, responseBody) @@ -54,18 +48,27 @@ func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http usage.PromptTokensDetails.CacheWriteTokens = responsesResponse.Usage.InputTokensDetails.CacheWriteTokens } } - if info == nil || info.ResponsesUsageInfo == nil || info.ResponsesUsageInfo.BuiltInTools == nil { - return &usage, nil - } - // 解析 Tools 用量 - for _, tool := range responsesResponse.Tools { - buildToolinfo, ok := info.ResponsesUsageInfo.BuiltInTools[common.Interface2String(tool["type"])] - if !ok || buildToolinfo == nil { - logger.LogError(c, fmt.Sprintf("BuiltInTools not found for tool type: %v", tool["type"])) - continue + // Count actual tool invocations from Output (not tool declarations). + for _, output := range responsesResponse.Output { + switch output.Type { + case dto.BuildInCallWebSearchCall: + info.CountBillableToolCall(dto.BuildInCallWebSearchCall, "") + case dto.BuildInCallFileSearchCall: + info.CountBillableToolCall(dto.BuildInCallFileSearchCall, "") + case dto.BuildInCallFunctionCall: + info.CountBillableToolCall(dto.BuildInCallFunctionCall, output.Name) } - buildToolinfo.CallCount++ } + + imageCounter := &relaycommon.ImageGenerationCallCounter{} + if !relaycommon.IsNonBillableResponsesStatus(responsesResponse.Status) { + for i := range responsesResponse.Output { + idx := i + imageCounter.Observe(&responsesResponse.Output[i], &idx) + } + } + imageCounter.Commit(info) + return &usage, nil } @@ -79,6 +82,8 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp var usage = &dto.Usage{} var responseTextBuilder strings.Builder + imageCounter := &relaycommon.ImageGenerationCallCounter{} + imageCommitted := false helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { @@ -91,7 +96,7 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp } sendResponsesStreamData(c, streamResponse, data) switch streamResponse.Type { - case "response.completed": + case "response.completed", "response.done": if streamResponse.Response != nil { if streamResponse.Response.Usage != nil { if streamResponse.Response.Usage.InputTokens != 0 { @@ -108,24 +113,45 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp usage.PromptTokensDetails.CacheWriteTokens = streamResponse.Response.Usage.InputTokensDetails.CacheWriteTokens } } - if streamResponse.Response.HasImageGenerationCall() { - c.Set("image_generation_call", true) - c.Set("image_generation_call_quality", streamResponse.Response.GetQuality()) - c.Set("image_generation_call_size", streamResponse.Response.GetSize()) + if !imageCommitted { + if relaycommon.IsNonBillableResponsesStatus(streamResponse.Response.Status) { + imageCounter.Reset() + imageCounter.Commit(info) + imageCommitted = true + } else { + for i := range streamResponse.Response.Output { + idx := i + imageCounter.Observe(&streamResponse.Response.Output[i], &idx) + } + imageCounter.Commit(info) + imageCommitted = true + } } + } else if !imageCommitted { + imageCounter.Commit(info) + imageCommitted = true + } + case "response.failed", "response.incomplete", "response.cancelled", "response.canceled": + if !imageCommitted { + imageCounter.Reset() + imageCounter.Commit(info) + imageCommitted = true } case "response.output_text.delta": // 处理输出文本 responseTextBuilder.WriteString(streamResponse.Delta) case dto.ResponsesOutputTypeItemDone: - // 函数调用处理 if streamResponse.Item != nil { switch streamResponse.Item.Type { case dto.BuildInCallWebSearchCall: - if info != nil && info.ResponsesUsageInfo != nil && info.ResponsesUsageInfo.BuiltInTools != nil { - if webSearchTool, exists := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview]; exists && webSearchTool != nil { - webSearchTool.CallCount++ - } + info.CountBillableToolCall(dto.BuildInCallWebSearchCall, "") + case dto.BuildInCallFileSearchCall: + info.CountBillableToolCall(dto.BuildInCallFileSearchCall, "") + case dto.BuildInCallFunctionCall: + info.CountBillableToolCall(dto.BuildInCallFunctionCall, streamResponse.Item.Name) + case dto.ResponsesOutputTypeImageGenerationCall: + if !imageCommitted { + imageCounter.Observe(streamResponse.Item, streamResponse.OutputIndex) } } } diff --git a/relay/channel/openai/relay_responses_billing_test.go b/relay/channel/openai/relay_responses_billing_test.go new file mode 100644 index 00000000..febe69c4 --- /dev/null +++ b/relay/channel/openai/relay_responses_billing_test.go @@ -0,0 +1,267 @@ +package openai + +import ( + "bytes" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/setting/operation_setting" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestOaiResponsesHandlerCountsOutputCallsNotDeclarations(t *testing.T) { + gin.SetMode(gin.TestMode) + operation_setting.SetToolPriceForTest("priced_fn", 5.0) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest("priced_fn") + }) + + body, err := common.Marshal(dto.OpenAIResponsesResponse{ + Tools: []map[string]any{ + {"type": "web_search_preview"}, + {"type": "file_search"}, + }, + Output: []dto.ResponsesOutput{ + {Type: dto.BuildInCallWebSearchCall}, + {Type: dto.BuildInCallWebSearchCall}, + {Type: dto.BuildInCallFunctionCall, Name: "priced_fn"}, + {Type: dto.BuildInCallFunctionCall, Name: "unpriced_fn"}, + }, + Usage: &dto.Usage{InputTokens: 1, OutputTokens: 1, TotalTokens: 2}, + }) + require.NoError(t, err) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + + info := &relaycommon.RelayInfo{ + OriginModelName: "gpt-5.1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolWebSearchPreview: {ToolName: dto.BuildInToolWebSearchPreview, CallCount: 0}, + dto.BuildInToolFileSearch: {ToolName: dto.BuildInToolFileSearch, CallCount: 0}, + }, + }, + } + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"Content-Type": []string{"application/json"}}, + } + + usage, apiErr := OaiResponsesHandler(c, info, resp) + require.Nil(t, apiErr) + require.NotNil(t, usage) + assert.Equal(t, 2, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview].CallCount) + assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolFileSearch].CallCount) + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, "priced_fn") + assert.Equal(t, 1, info.ResponsesUsageInfo.BuiltInTools["priced_fn"].CallCount) + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, "unpriced_fn") +} + +func TestOaiResponsesHandlerDeclaredToolsWithoutOutputCountZero(t *testing.T) { + gin.SetMode(gin.TestMode) + + body, err := common.Marshal(dto.OpenAIResponsesResponse{ + Tools: []map[string]any{ + {"type": "web_search_preview"}, + {"type": "file_search"}, + }, + Output: []dto.ResponsesOutput{ + {Type: "message", Role: "assistant"}, + }, + Usage: &dto.Usage{InputTokens: 1, OutputTokens: 1, TotalTokens: 2}, + }) + require.NoError(t, err) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + + info := &relaycommon.RelayInfo{ + OriginModelName: "gpt-5.1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolWebSearchPreview: {ToolName: dto.BuildInToolWebSearchPreview, CallCount: 0}, + dto.BuildInToolFileSearch: {ToolName: dto.BuildInToolFileSearch, CallCount: 0}, + }, + }, + } + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"Content-Type": []string{"application/json"}}, + } + + _, apiErr := OaiResponsesHandler(c, info, resp) + require.Nil(t, apiErr) + assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview].CallCount) + assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolFileSearch].CallCount) +} + +func TestOaiResponsesHandlerCountsCompletedImageGenerationOutputs(t *testing.T) { + gin.SetMode(gin.TestMode) + + body, err := common.Marshal(dto.OpenAIResponsesResponse{ + Status: []byte(`"completed"`), + Output: []dto.ResponsesOutput{ + { + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "completed", + Result: "base64-a", + }, + { + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_2", + Status: "completed", + Result: "base64-b", + }, + { + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_empty", + Status: "completed", + Result: "", + }, + }, + Usage: &dto.Usage{InputTokens: 1, OutputTokens: 1, TotalTokens: 2}, + }) + require.NoError(t, err) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + info := &relaycommon.RelayInfo{OriginModelName: "gpt-5.1"} + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"Content-Type": []string{"application/json"}}, + } + + _, apiErr := OaiResponsesHandler(c, info, resp) + require.Nil(t, apiErr) + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolImageGeneration) + assert.Equal(t, 2, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount) + assert.False(t, c.GetBool("image_generation_call")) +} + +func TestOaiResponsesHandlerIncompleteStatusCommitsZeroImageGeneration(t *testing.T) { + gin.SetMode(gin.TestMode) + + body, err := common.Marshal(dto.OpenAIResponsesResponse{ + Status: []byte(`"incomplete"`), + Output: []dto.ResponsesOutput{ + { + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "completed", + Result: "base64-a", + }, + }, + Usage: &dto.Usage{InputTokens: 1, OutputTokens: 1, TotalTokens: 2}, + }) + require.NoError(t, err) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + info := &relaycommon.RelayInfo{ + OriginModelName: "gpt-5.1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolImageGeneration: {ToolName: dto.BuildInToolImageGeneration, CallCount: 0}, + }, + }, + } + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(body)), + Header: http.Header{"Content-Type": []string{"application/json"}}, + } + + _, apiErr := OaiResponsesHandler(c, info, resp) + require.Nil(t, apiErr) + assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount) +} + +func runResponsesImageBillingStream(t *testing.T, events ...string) *relaycommon.RelayInfo { + t.Helper() + gin.SetMode(gin.TestMode) + oldTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 30 + t.Cleanup(func() { + constant.StreamingTimeout = oldTimeout + }) + + var body strings.Builder + for _, event := range events { + body.WriteString("data: ") + body.WriteString(event) + body.WriteString("\n\n") + } + body.WriteString("data: [DONE]\n\n") + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + c.Set(common.RequestIdKey, "responses-image-billing-test") + info := &relaycommon.RelayInfo{ + OriginModelName: "gpt-5.1", + DisablePing: true, + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: "gpt-5.1", + }, + } + resp := &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(body.String())), + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + } + + _, apiErr := OaiResponsesStreamHandler(c, info, resp) + require.Nil(t, apiErr) + require.NotNil(t, info.ResponsesUsageInfo) + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolImageGeneration) + return info +} + +func TestOaiResponsesStreamHandlerDeduplicatesCompletedImageOutput(t *testing.T) { + item := `{"type":"image_generation_call","id":"img_1","call_id":"call_1","status":"completed","result":"base64-a"}` + info := runResponsesImageBillingStream( + t, + `{"type":"response.output_item.done","output_index":0,"item":`+item+`}`, + `{"type":"response.completed","response":{"status":"completed","output":[`+item+`],"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}`, + ) + + assert.Equal(t, 1, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount) +} + +func TestOaiResponsesStreamHandlerDiscardsImageOutputOnIncomplete(t *testing.T) { + info := runResponsesImageBillingStream( + t, + `{"type":"response.output_item.done","output_index":0,"item":{"type":"image_generation_call","id":"img_1","status":"completed","result":"base64-a"}}`, + `{"type":"response.incomplete","response":{"status":"incomplete"}}`, + ) + + assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount) +} + +func TestOaiResponsesStreamHandlerDoesNotCountPartialImageEvent(t *testing.T) { + info := runResponsesImageBillingStream( + t, + `{"type":"response.image_generation_call.partial_image","output_index":0,"partial_image_b64":"partial-bytes"}`, + `{"type":"response.completed","response":{"status":"completed","output":[]}}`, + ) + + assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount) +} diff --git a/relay/channel/openai/stream_tool_billing_test.go b/relay/channel/openai/stream_tool_billing_test.go new file mode 100644 index 00000000..c7922363 --- /dev/null +++ b/relay/channel/openai/stream_tool_billing_test.go @@ -0,0 +1,27 @@ +package openai + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestCollectStreamFunctionCallNamesDedupesSameIndex(t *testing.T) { + seen := make(map[string]struct{}) + var names []string + + chunks := []string{ + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"c1","type":"function","function":{"name":"get_weather","arguments":""}}]}}]}`, + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"q\":"}}]}}]}`, + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"x\"}"}}]}}]}`, + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":1,"id":"c2","type":"function","function":{"name":"get_time","arguments":""}}]}}]}`, + `{"choices":[{"index":0,"delta":{"tool_calls":[{"index":1,"function":{"arguments":"{}"}}]}}]}`, + } + for _, chunk := range chunks { + collectStreamFunctionCallNames(chunk, seen, &names) + } + + require.Len(t, names, 2) + assert.Equal(t, []string{"get_weather", "get_time"}, names) +} diff --git a/relay/channel/sub2api/adaptor.go b/relay/channel/sub2api/adaptor.go new file mode 100644 index 00000000..fa454f2a --- /dev/null +++ b/relay/channel/sub2api/adaptor.go @@ -0,0 +1,121 @@ +package sub2api + +import ( + "errors" + "io" + "net/http" + + "github.com/QuantumNous/new-api/dto" + "github.com/QuantumNous/new-api/relay/channel" + "github.com/QuantumNous/new-api/relay/channel/claude" + "github.com/QuantumNous/new-api/relay/channel/gemini" + "github.com/QuantumNous/new-api/relay/channel/openai" + relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/QuantumNous/new-api/types" + + "github.com/gin-gonic/gin" +) + +type Adaptor struct { + openaiAdaptor openai.Adaptor + claudeAdaptor claude.Adaptor + geminiAdaptor gemini.Adaptor +} + +func (a *Adaptor) Init(info *relaycommon.RelayInfo) { + a.openaiAdaptor.Init(info) + a.claudeAdaptor.Init(info) + a.geminiAdaptor.Init(info) +} + +func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { + if info.RelayMode == relayconstant.RelayModeAlphaSearch { + return relaycommon.GetFullRequestURL(info.ChannelBaseUrl, "/v1/alpha/search", info.ChannelType), nil + } + return relaycommon.GetFullRequestURL(info.ChannelBaseUrl, info.RequestURLPath, info.ChannelType), nil +} + +func (a *Adaptor) SetupRequestHeader(c *gin.Context, req *http.Header, info *relaycommon.RelayInfo) error { + channel.SetupApiRequestHeader(info, c, req) + req.Set("Authorization", "Bearer "+info.ApiKey) + + switch info.RelayFormat { + case types.RelayFormatClaude: + req.Set("x-api-key", info.ApiKey) + if req.Get("anthropic-version") == "" { + anthropicVersion := c.Request.Header.Get("anthropic-version") + if anthropicVersion == "" { + anthropicVersion = "2023-06-01" + } + req.Set("anthropic-version", anthropicVersion) + } + case types.RelayFormatGemini: + req.Set("x-goog-api-key", info.ApiKey) + } + return nil +} + +func (a *Adaptor) ConvertOpenAIRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeneralOpenAIRequest) (any, error) { + if request == nil { + return nil, errors.New("request is nil") + } + return request, nil +} + +func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) { + return request, nil +} + +func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.EmbeddingRequest) (any, error) { + return request, nil +} + +func (a *Adaptor) ConvertClaudeRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.ClaudeRequest) (any, error) { + if request == nil { + return nil, errors.New("request is nil") + } + return request, nil +} + +func (a *Adaptor) ConvertGeminiRequest(c *gin.Context, info *relaycommon.RelayInfo, request *dto.GeminiChatRequest) (any, error) { + if request == nil { + return nil, errors.New("request is nil") + } + return request, nil +} + +func (a *Adaptor) ConvertImageRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.ImageRequest) (any, error) { + return request, nil +} + +func (a *Adaptor) ConvertRerankRequest(c *gin.Context, relayMode int, request dto.RerankRequest) (any, error) { + return nil, errors.New("endpoint not supported") +} + +func (a *Adaptor) ConvertAudioRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.AudioRequest) (io.Reader, error) { + return nil, errors.New("endpoint not supported") +} + +func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error) { + return channel.DoApiRequest(a, c, info, requestBody) +} + +func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) { + switch info.RelayFormat { + case types.RelayFormatClaude: + return a.claudeAdaptor.DoResponse(c, resp, info) + case types.RelayFormatGemini: + return a.geminiAdaptor.DoResponse(c, resp, info) + default: + return a.openaiAdaptor.DoResponse(c, resp, info) + } +} + +func (a *Adaptor) GetModelList() []string { + return ModelList +} + +func (a *Adaptor) GetChannelName() string { + return ChannelName +} diff --git a/relay/channel/sub2api/adaptor_test.go b/relay/channel/sub2api/adaptor_test.go new file mode 100644 index 00000000..ed831963 --- /dev/null +++ b/relay/channel/sub2api/adaptor_test.go @@ -0,0 +1,27 @@ +package sub2api + +import ( + "testing" + + "github.com/QuantumNous/new-api/constant" + relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGetRequestURLAlphaSearch(t *testing.T) { + adaptor := &Adaptor{} + info := &relaycommon.RelayInfo{ + ChannelMeta: &relaycommon.ChannelMeta{ + ChannelType: constant.ChannelTypeSub2API, + ChannelBaseUrl: "https://sub2api.example", + }, + RequestURLPath: "/v1/alpha/search", + RelayMode: relayconstant.RelayModeAlphaSearch, + } + + url, err := adaptor.GetRequestURL(info) + require.NoError(t, err) + assert.Equal(t, "https://sub2api.example/v1/alpha/search", url) +} diff --git a/relay/channel/sub2api/constants.go b/relay/channel/sub2api/constants.go new file mode 100644 index 00000000..e18b12e2 --- /dev/null +++ b/relay/channel/sub2api/constants.go @@ -0,0 +1,6 @@ +package sub2api + +const ChannelName = "sub2api" + +// ModelList is empty because models are fetched dynamically from upstream /v1/models. +var ModelList = []string{} diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index 128e52e0..81df500f 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -342,6 +342,7 @@ var streamSupportedChannels = map[int]bool{ constant.ChannelTypeMiniMax: true, constant.ChannelTypeSiliconFlow: true, constant.ChannelTypeAdvancedCustom: true, + constant.ChannelTypeSub2API: true, constant.ChannelTypeTencent: true, } @@ -576,6 +577,11 @@ func GenRelayInfo(c *gin.Context, relayFormat types.RelayFormat, request dto.Req return GenRelayInfoResponsesCompaction(c, request), nil } return nil, errors.New("request is not a OpenAIResponsesCompactionRequest") + case types.RelayFormatOpenAIAlphaSearch: + if request, ok := request.(*dto.AlphaSearchRequest); ok { + return GenRelayInfoAlphaSearch(c, request), nil + } + return nil, errors.New("request is not a AlphaSearchRequest") case types.RelayFormatTask: info = genBaseRelayInfo(c, nil) info.TaskRelayInfo = &TaskRelayInfo{} @@ -650,6 +656,23 @@ func GenRelayInfoResponsesCompaction(c *gin.Context, request *dto.OpenAIResponse return info } +func GenRelayInfoAlphaSearch(c *gin.Context, request *dto.AlphaSearchRequest) *RelayInfo { + info := genBaseRelayInfo(c, request) + if info.RelayMode == relayconstant.RelayModeUnknown { + info.RelayMode = relayconstant.RelayModeAlphaSearch + } + info.RelayFormat = types.RelayFormatOpenAIAlphaSearch + info.ResponsesUsageInfo = &ResponsesUsageInfo{ + BuiltInTools: map[string]*BuildInToolInfo{ + dto.BuildInToolWebSearchPreview: { + ToolName: dto.BuildInToolWebSearchPreview, + CallCount: 0, + }, + }, + } + return info +} + //func (info *RelayInfo) SetPromptTokens(promptTokens int) { // info.promptTokens = promptTokens //} diff --git a/relay/common/tool_usage.go b/relay/common/tool_usage.go new file mode 100644 index 00000000..c53c5c54 --- /dev/null +++ b/relay/common/tool_usage.go @@ -0,0 +1,194 @@ +package common + +import ( + "crypto/sha256" + "encoding/hex" + "fmt" + "strings" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/dto" + "github.com/QuantumNous/new-api/setting/operation_setting" +) + +var reservedBillableToolNames = map[string]struct{}{ + dto.BuildInToolWebSearchPreview: {}, + dto.BuildInToolWebSearch: {}, + dto.BuildInToolFileSearch: {}, + dto.BuildInToolGoogleSearch: {}, + dto.BuildInToolImageGeneration: {}, +} + +// CountBillableToolCall is the single entry point for per-call tool billing counts. +// Built-in call types always count; custom function/tool_use names only count when priced. +func (info *RelayInfo) CountBillableToolCall(itemType string, functionName string) { + if info == nil { + return + } + if info.ResponsesUsageInfo == nil { + info.ResponsesUsageInfo = &ResponsesUsageInfo{ + BuiltInTools: make(map[string]*BuildInToolInfo), + } + } + if info.ResponsesUsageInfo.BuiltInTools == nil { + info.ResponsesUsageInfo.BuiltInTools = make(map[string]*BuildInToolInfo) + } + + switch itemType { + case dto.BuildInCallWebSearchCall: + info.incrementBillableToolCall(resolveWebSearchToolName(info.ResponsesUsageInfo.BuiltInTools)) + case dto.BuildInCallFileSearchCall: + info.incrementBillableToolCall(dto.BuildInToolFileSearch) + case dto.BuildInCallFunctionCall, dto.BuildInCallToolUse: + if functionName == "" { + return + } + if _, reserved := reservedBillableToolNames[functionName]; reserved { + return + } + if operation_setting.GetToolPriceForModel(functionName, info.OriginModelName) <= 0 { + return + } + info.incrementBillableToolCall(functionName) + } +} + +func resolveWebSearchToolName(tools map[string]*BuildInToolInfo) string { + if _, ok := tools[dto.BuildInToolWebSearchPreview]; ok { + return dto.BuildInToolWebSearchPreview + } + if _, ok := tools[dto.BuildInToolWebSearch]; ok { + return dto.BuildInToolWebSearch + } + return dto.BuildInToolWebSearchPreview +} + +func (info *RelayInfo) incrementBillableToolCall(name string) { + if existing, ok := info.ResponsesUsageInfo.BuiltInTools[name]; ok && existing != nil { + existing.CallCount++ + return + } + info.ResponsesUsageInfo.BuiltInTools[name] = &BuildInToolInfo{ + ToolName: name, + CallCount: 1, + } +} + +// ImageGenerationCallCounter counts completed Responses image_generation_call +// outputs with stream-safe identity deduplication. +type ImageGenerationCallCounter struct { + seen map[string]struct{} + count int +} + +// Observe records one completed final image output when billable. +// outputIndex may be nil; when set and nonnegative it participates in dedup. +func (c *ImageGenerationCallCounter) Observe(item *dto.ResponsesOutput, outputIndex *int) { + if c == nil || item == nil { + return + } + if item.Type != dto.ResponsesOutputTypeImageGenerationCall { + return + } + if strings.TrimSpace(item.Result) == "" { + return + } + switch strings.ToLower(strings.TrimSpace(item.Status)) { + case "failed", "cancelled", "canceled", "incomplete", "partial": + return + } + + aliases := make([]string, 0, 4) + if item.ID != "" { + aliases = append(aliases, "id:"+item.ID) + } + if item.CallId != "" { + aliases = append(aliases, "call:"+item.CallId) + } + if outputIndex != nil && *outputIndex >= 0 { + aliases = append(aliases, fmt.Sprintf("index:%d", *outputIndex)) + } + sum := sha256.Sum256([]byte(item.Result)) + aliases = append(aliases, "result:"+hex.EncodeToString(sum[:])) + + if c.seen == nil { + c.seen = make(map[string]struct{}) + } + for _, alias := range aliases { + if _, ok := c.seen[alias]; ok { + return + } + } + for _, alias := range aliases { + c.seen[alias] = struct{}{} + } + c.count++ +} + +// Reset clears pending observations (used when a terminal response fails). +func (c *ImageGenerationCallCounter) Reset() { + if c == nil { + return + } + c.seen = nil + c.count = 0 +} + +// Count returns the deduplicated completed image output count before commit capping. +func (c *ImageGenerationCallCounter) Count() int { + if c == nil { + return 0 + } + return c.count +} + +// Commit writes the capped completed-output count into RelayInfo once. +// Request tool declarations alone must not become billable calls. +func (c *ImageGenerationCallCounter) Commit(info *RelayInfo) { + if info == nil { + return + } + if info.ResponsesUsageInfo == nil { + info.ResponsesUsageInfo = &ResponsesUsageInfo{ + BuiltInTools: make(map[string]*BuildInToolInfo), + } + } + if info.ResponsesUsageInfo.BuiltInTools == nil { + info.ResponsesUsageInfo.BuiltInTools = make(map[string]*BuildInToolInfo) + } + + count := 0 + if c != nil { + count = c.count + } + if count > dto.MaxImageN { + count = dto.MaxImageN + } + + if existing, ok := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]; ok && existing != nil { + existing.CallCount = count + return + } + info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration] = &BuildInToolInfo{ + ToolName: dto.BuildInToolImageGeneration, + CallCount: count, + } +} + +// IsNonBillableResponsesStatus reports terminal response statuses that must not +// bill pending image_generation observations. +func IsNonBillableResponsesStatus(status []byte) bool { + if len(status) == 0 { + return false + } + var s string + if err := common.Unmarshal(status, &s); err != nil { + return false + } + switch strings.ToLower(strings.TrimSpace(s)) { + case "failed", "cancelled", "canceled", "incomplete": + return true + default: + return false + } +} diff --git a/relay/common/tool_usage_test.go b/relay/common/tool_usage_test.go new file mode 100644 index 00000000..1ef92d7c --- /dev/null +++ b/relay/common/tool_usage_test.go @@ -0,0 +1,348 @@ +package common + +import ( + "strings" + "testing" + + "github.com/QuantumNous/new-api/dto" + "github.com/QuantumNous/new-api/setting/operation_setting" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestCountBillableToolCallWebSearchPrefersDeclaredWebSearch(t *testing.T) { + info := &RelayInfo{ + OriginModelName: "gpt-5.1", + ResponsesUsageInfo: &ResponsesUsageInfo{ + BuiltInTools: map[string]*BuildInToolInfo{ + dto.BuildInToolWebSearch: {ToolName: dto.BuildInToolWebSearch, CallCount: 0}, + }, + }, + } + + info.CountBillableToolCall(dto.BuildInCallWebSearchCall, "") + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolWebSearch) + assert.Equal(t, 1, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearch].CallCount) + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolWebSearchPreview) +} + +func TestCountBillableToolCallWebSearchDefaultsToPreview(t *testing.T) { + info := &RelayInfo{OriginModelName: "gpt-5.1"} + + info.CountBillableToolCall(dto.BuildInCallWebSearchCall, "") + require.NotNil(t, info.ResponsesUsageInfo) + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolWebSearchPreview) + assert.Equal(t, 1, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview].CallCount) +} + +func TestCountBillableToolCallFunctionCallRequiresPrice(t *testing.T) { + operation_setting.SetToolPriceForTest("my_priced_fn", 5.0) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest("my_priced_fn") + }) + + info := &RelayInfo{OriginModelName: "gpt-5.1"} + info.CountBillableToolCall(dto.BuildInCallFunctionCall, "my_priced_fn") + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, "my_priced_fn") + assert.Equal(t, 1, info.ResponsesUsageInfo.BuiltInTools["my_priced_fn"].CallCount) + + info.CountBillableToolCall(dto.BuildInCallFunctionCall, "unpriced_fn") + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, "unpriced_fn") +} + +func TestCountBillableToolCallFunctionCallSkipsReservedNames(t *testing.T) { + info := &RelayInfo{OriginModelName: "gpt-5.1"} + + info.CountBillableToolCall(dto.BuildInCallFunctionCall, dto.BuildInToolWebSearchPreview) + info.CountBillableToolCall(dto.BuildInCallFunctionCall, dto.BuildInToolFileSearch) + info.CountBillableToolCall(dto.BuildInCallFunctionCall, dto.BuildInToolGoogleSearch) + info.CountBillableToolCall(dto.BuildInCallFunctionCall, dto.BuildInToolImageGeneration) + + if info.ResponsesUsageInfo != nil { + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolWebSearchPreview) + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolFileSearch) + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolGoogleSearch) + assert.NotContains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolImageGeneration) + } +} + +func TestImageGenerationCallCounterCompletedOutputs(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + observe func(c *ImageGenerationCallCounter) + wantCount int + }{ + { + name: "one final result", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "completed", + Result: "base64-a", + }, &idx) + }, + wantCount: 1, + }, + { + name: "two distinct finals", + observe: func(c *ImageGenerationCallCounter) { + idx0, idx1 := 0, 1 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Result: "base64-a", + }, &idx0) + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_2", + Result: "base64-b", + }, &idx1) + }, + wantCount: 2, + }, + { + name: "empty result", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Result: " ", + }, &idx) + }, + wantCount: 0, + }, + { + name: "failed status", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "failed", + Result: "base64-a", + }, &idx) + }, + wantCount: 0, + }, + { + name: "incomplete status", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "incomplete", + Result: "base64-a", + }, &idx) + }, + wantCount: 0, + }, + { + name: "cancelled status", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "cancelled", + Result: "base64-a", + }, &idx) + }, + wantCount: 0, + }, + { + name: "canceled status", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "canceled", + Result: "base64-a", + }, &idx) + }, + wantCount: 0, + }, + { + name: "partial status", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "partial", + Result: "partial-bytes", + }, &idx) + }, + wantCount: 0, + }, + { + name: "id dedup", + observe: func(c *ImageGenerationCallCounter) { + idx0, idx1 := 0, 1 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + CallId: "call_a", + Result: "base64-a", + }, &idx0) + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + CallId: "call_b", + Result: "base64-b", + }, &idx1) + }, + wantCount: 1, + }, + { + name: "index dedup", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Result: "base64-a", + }, &idx) + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_2", + Result: "base64-b", + }, &idx) + }, + wantCount: 1, + }, + { + name: "result hash dedup", + observe: func(c *ImageGenerationCallCounter) { + idx0, idx1 := 0, 1 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + Result: "same-bytes", + }, &idx0) + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + Result: "same-bytes", + }, &idx1) + }, + wantCount: 1, + }, + { + name: "output_item.done plus completed dedup", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + item := &dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + CallId: "call_1", + Status: "completed", + Result: "base64-a", + } + c.Observe(item, &idx) + c.Observe(item, &idx) + }, + wantCount: 1, + }, + { + name: "output_item.done plus incomplete equals zero", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "completed", + Result: "base64-a", + }, &idx) + c.Reset() + }, + wantCount: 0, + }, + { + name: "partial event equals zero", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: "image_generation_call.partial_image", + ID: "img_1", + Result: "partial-bytes", + }, &idx) + }, + wantCount: 0, + }, + { + name: "in_progress with final result counts", + observe: func(c *ImageGenerationCallCounter) { + idx := 0 + c.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_1", + Status: "in_progress", + Result: "base64-a", + }, &idx) + }, + wantCount: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + counter := &ImageGenerationCallCounter{} + tt.observe(counter) + assert.Equal(t, tt.wantCount, counter.Count()) + }) + } +} + +func TestImageGenerationCallCounterCommitCapsAtMaxImageN(t *testing.T) { + t.Parallel() + + counter := &ImageGenerationCallCounter{} + for i := 0; i < dto.MaxImageN+3; i++ { + idx := i + counter.Observe(&dto.ResponsesOutput{ + Type: dto.ResponsesOutputTypeImageGenerationCall, + ID: "img_" + strings.Repeat("a", i+1), + Result: "result-" + strings.Repeat("b", i+1), + }, &idx) + } + require.Equal(t, dto.MaxImageN+3, counter.Count()) + + info := &RelayInfo{} + counter.Commit(info) + require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolImageGeneration) + assert.Equal(t, dto.MaxImageN, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount) +} + +func TestImageGenerationCallCounterCommitDoesNotBillDeclarationsAlone(t *testing.T) { + t.Parallel() + + info := &RelayInfo{ + ResponsesUsageInfo: &ResponsesUsageInfo{ + BuiltInTools: map[string]*BuildInToolInfo{ + dto.BuildInToolImageGeneration: { + ToolName: dto.BuildInToolImageGeneration, + CallCount: 0, + }, + }, + }, + } + (&ImageGenerationCallCounter{}).Commit(info) + assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount) +} + +func TestIsNonBillableResponsesStatus(t *testing.T) { + t.Parallel() + + assert.True(t, IsNonBillableResponsesStatus([]byte(`"failed"`))) + assert.True(t, IsNonBillableResponsesStatus([]byte(`"incomplete"`))) + assert.True(t, IsNonBillableResponsesStatus([]byte(`"cancelled"`))) + assert.True(t, IsNonBillableResponsesStatus([]byte(`"canceled"`))) + assert.False(t, IsNonBillableResponsesStatus([]byte(`"completed"`))) + assert.False(t, IsNonBillableResponsesStatus(nil)) +} diff --git a/relay/constant/relay_mode.go b/relay/constant/relay_mode.go index 25671567..5f4b3be5 100644 --- a/relay/constant/relay_mode.go +++ b/relay/constant/relay_mode.go @@ -52,6 +52,8 @@ const ( RelayModeGemini RelayModeResponsesCompact + + RelayModeAlphaSearch ) func Path2RelayMode(path string) int { @@ -76,6 +78,8 @@ func Path2RelayMode(path string) int { relayMode = RelayModeResponsesCompact } else if strings.HasPrefix(path, "/v1/responses") { relayMode = RelayModeResponses + } else if strings.HasPrefix(path, "/v1/alpha/search") { + relayMode = RelayModeAlphaSearch } else if strings.HasPrefix(path, "/v1/audio/speech") { relayMode = RelayModeAudioSpeech } else if strings.HasPrefix(path, "/v1/audio/transcriptions") { diff --git a/relay/constant/relay_mode_test.go b/relay/constant/relay_mode_test.go new file mode 100644 index 00000000..6c419815 --- /dev/null +++ b/relay/constant/relay_mode_test.go @@ -0,0 +1,22 @@ +package constant + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestPath2RelayMode(t *testing.T) { + tests := []struct { + path string + want int + }{ + {path: "/v1/alpha/search", want: RelayModeAlphaSearch}, + {path: "/v1/alpha/search?foo=1", want: RelayModeAlphaSearch}, + } + for _, tt := range tests { + t.Run(tt.path, func(t *testing.T) { + assert.Equal(t, tt.want, Path2RelayMode(tt.path)) + }) + } +} diff --git a/relay/helper/valid_request.go b/relay/helper/valid_request.go index afbd59ff..22068a45 100644 --- a/relay/helper/valid_request.go +++ b/relay/helper/valid_request.go @@ -38,6 +38,8 @@ func GetAndValidateRequest(c *gin.Context, format types.RelayFormat) (request dt request, err = GetAndValidateResponsesRequest(c) case types.RelayFormatOpenAIResponsesCompaction: request, err = GetAndValidateResponsesCompactionRequest(c) + case types.RelayFormatOpenAIAlphaSearch: + request, err = GetAndValidateAlphaSearchRequest(c) case types.RelayFormatOpenAIImage: request, err = GetAndValidOpenAIImageRequest(c, relayMode) @@ -146,6 +148,26 @@ func GetAndValidateResponsesRequest(c *gin.Context) (*dto.OpenAIResponsesRequest return request, nil } +func GetAndValidateAlphaSearchRequest(c *gin.Context) (*dto.AlphaSearchRequest, error) { + request := &dto.AlphaSearchRequest{} + if err := common.UnmarshalBodyReusable(c, request); err != nil { + return nil, err + } + if request.Model == "" { + return nil, errors.New("model is required") + } + storage, err := common.GetBodyStorage(c) + if err != nil { + return nil, err + } + rawBody, err := storage.Bytes() + if err != nil { + return nil, err + } + request.RawBody = rawBody + return request, nil +} + func GetAndValidateResponsesCompactionRequest(c *gin.Context) (*dto.OpenAIResponsesCompactionRequest, error) { request := &dto.OpenAIResponsesCompactionRequest{} if err := common.UnmarshalBodyReusable(c, request); err != nil { diff --git a/relay/relay_adaptor.go b/relay/relay_adaptor.go index 8920ba8c..9f768e4b 100644 --- a/relay/relay_adaptor.go +++ b/relay/relay_adaptor.go @@ -30,6 +30,7 @@ import ( "github.com/QuantumNous/new-api/relay/channel/perplexity" "github.com/QuantumNous/new-api/relay/channel/replicate" "github.com/QuantumNous/new-api/relay/channel/siliconflow" + "github.com/QuantumNous/new-api/relay/channel/sub2api" "github.com/QuantumNous/new-api/relay/channel/submodel" taskali "github.com/QuantumNous/new-api/relay/channel/task/ali" taskdoubao "github.com/QuantumNous/new-api/relay/channel/task/doubao" @@ -123,6 +124,8 @@ func GetAdaptor(apiType int) channel.Adaptor { return &codex.Adaptor{} case constant.APITypeAdvancedCustom: return &advancedcustom.Adaptor{} + case constant.APITypeSub2API: + return &sub2api.Adaptor{} } return nil } diff --git a/router/relay-router.go b/router/relay-router.go index 17a13cad..5fa64ef5 100644 --- a/router/relay-router.go +++ b/router/relay-router.go @@ -105,6 +105,11 @@ func SetRelayRouter(router *gin.Engine) { controller.Relay(c, types.RelayFormatOpenAIResponsesCompaction) }) + // alpha search related routes (Codex standalone web search) + httpRouter.POST("/alpha/search", func(c *gin.Context) { + controller.Relay(c, types.RelayFormatOpenAIAlphaSearch) + }) + // image related routes httpRouter.POST("/edits", func(c *gin.Context) { controller.Relay(c, types.RelayFormatOpenAIImage) diff --git a/service/text_quota.go b/service/text_quota.go index 7da33912..414a98c1 100644 --- a/service/text_quota.go +++ b/service/text_quota.go @@ -2,6 +2,8 @@ package service import ( "fmt" + "math" + "sort" "strings" "time" @@ -13,6 +15,7 @@ import ( "github.com/QuantumNous/new-api/pkg/billingexpr" perfmetrics "github.com/QuantumNous/new-api/pkg/perf_metrics" relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/types" @@ -21,40 +24,56 @@ import ( "github.com/shopspring/decimal" ) +// ToolSurchargeItem is one billable tool-call line for consume logs. +type ToolSurchargeItem struct { + Name string `json:"name"` + Count int `json:"count"` + Price float64 `json:"price"` +} + +func appendToolSurchargeLogInfo(other map[string]interface{}, items []ToolSurchargeItem) { + if len(items) == 0 { + return + } + other["tool_surcharges"] = items +} + type textQuotaSummary struct { - PromptTokens int - CompletionTokens int - TotalTokens int - CacheTokens int - CacheCreationTokens int - CacheCreationTokens5m int - CacheCreationTokens1h int - ImageTokens int - AudioTokens int - ModelName string - TokenName string - UseTimeSeconds int64 - CompletionRatio float64 - CacheRatio float64 - ImageRatio float64 - ModelRatio float64 - GroupRatio float64 - ModelPrice float64 - CacheCreationRatio float64 - CacheCreationRatio5m float64 - CacheCreationRatio1h float64 - Quota int - IsClaudeUsageSemantic bool - UsageSemantic string - WebSearchPrice float64 - WebSearchCallCount int - ClaudeWebSearchPrice float64 - ClaudeWebSearchCallCount int - FileSearchPrice float64 - FileSearchCallCount int - AudioInputPrice float64 - ImageGenerationCallPrice float64 - ToolCallSurchargeQuota decimal.Decimal + PromptTokens int + CompletionTokens int + TotalTokens int + CacheTokens int + CacheCreationTokens int + CacheCreationTokens5m int + CacheCreationTokens1h int + ImageTokens int + AudioTokens int + ModelName string + TokenName string + UseTimeSeconds int64 + CompletionRatio float64 + CacheRatio float64 + ImageRatio float64 + ModelRatio float64 + GroupRatio float64 + ModelPrice float64 + CacheCreationRatio float64 + CacheCreationRatio5m float64 + CacheCreationRatio1h float64 + Quota int + IsClaudeUsageSemantic bool + UsageSemantic string + AudioInputPrice float64 + ToolSurchargeItems []ToolSurchargeItem + ToolCallSurchargeQuota decimal.Decimal +} + +// hasBillableUsage reports whether this request should incur any charge. +// A request can carry zero tokens yet still be billable via a tool-call +// surcharge (e.g. /v1/alpha/search returns no usage but bills one web_search +// call), so token count alone is not sufficient to decide. +func (s *textQuotaSummary) hasBillableUsage() bool { + return s.TotalTokens > 0 || !s.ToolCallSurchargeQuota.IsZero() } func cacheWriteTokensTotal(summary textQuotaSummary) int { @@ -81,60 +100,91 @@ func isLegacyClaudeDerivedOpenAIUsage(relayInfo *relaycommon.RelayInfo, usage *d return usage.ClaudeCacheCreation5mTokens > 0 || usage.ClaudeCacheCreation1hTokens > 0 } +func collectToolSurchargeItem(items []ToolSurchargeItem, name string, count int, modelName string) []ToolSurchargeItem { + if count <= 0 { + return items + } + price := operation_setting.GetToolPriceForModel(name, modelName) + if price <= 0 || math.IsNaN(price) || math.IsInf(price, 0) { + return items + } + return append(items, ToolSurchargeItem{ + Name: name, + Count: count, + Price: price, + }) +} + +func mergeToolSurchargeItems(items []ToolSurchargeItem) []ToolSurchargeItem { + if len(items) == 0 { + return nil + } + sort.Slice(items, func(i, j int) bool { + if items[i].Name == items[j].Name { + return items[i].Price < items[j].Price + } + return items[i].Name < items[j].Name + }) + + merged := items[:0] + for _, item := range items { + lastIndex := len(merged) - 1 + if lastIndex >= 0 && + merged[lastIndex].Name == item.Name && + merged[lastIndex].Price == item.Price { + if item.Count > math.MaxInt-merged[lastIndex].Count { + common.SysError("tool surcharge call count overflow for " + item.Name) + merged[lastIndex].Count = math.MaxInt + } else { + merged[lastIndex].Count += item.Count + } + continue + } + merged = append(merged, item) + } + return merged +} + func calculateTextToolCallSurcharge(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, summary *textQuotaSummary) decimal.Decimal { dGroupRatio := decimal.NewFromFloat(summary.GroupRatio) dQuotaPerUnit := decimal.NewFromFloat(common.QuotaPerUnit) + var items []ToolSurchargeItem + + if relayInfo.ResponsesUsageInfo != nil { + for name, tool := range relayInfo.ResponsesUsageInfo.BuiltInTools { + if tool == nil { + continue + } + items = collectToolSurchargeItem(items, name, tool.CallCount, summary.ModelName) + } + } + if relayInfo.RelayMode != relayconstant.RelayModeResponses && + strings.HasSuffix(summary.ModelName, "search-preview") { + items = collectToolSurchargeItem(items, dto.BuildInToolWebSearchPreview, 1, summary.ModelName) + } + + items = collectToolSurchargeItem( + items, + dto.BuildInToolWebSearch, + ctx.GetInt("claude_web_search_requests"), + summary.ModelName, + ) + + if ctx.GetBool("gemini_google_search_call") { + items = collectToolSurchargeItem(items, dto.BuildInToolGoogleSearch, 1, summary.ModelName) + } + + summary.ToolSurchargeItems = mergeToolSurchargeItems(items) var surcharge decimal.Decimal - - if relayInfo.ResponsesUsageInfo != nil { - if webSearchTool, exists := relayInfo.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolWebSearchPreview]; exists && webSearchTool.CallCount > 0 { - summary.WebSearchCallCount = webSearchTool.CallCount - summary.WebSearchPrice = operation_setting.GetToolPriceForModel("web_search_preview", summary.ModelName) - surcharge = surcharge.Add(decimal.NewFromFloat(summary.WebSearchPrice). - Mul(decimal.NewFromInt(int64(webSearchTool.CallCount))). - Div(decimal.NewFromInt(1000)). - Mul(dGroupRatio). - Mul(dQuotaPerUnit)) - } - } else if strings.HasSuffix(summary.ModelName, "search-preview") { - summary.WebSearchCallCount = 1 - summary.WebSearchPrice = operation_setting.GetToolPriceForModel("web_search_preview", summary.ModelName) - surcharge = surcharge.Add(decimal.NewFromFloat(summary.WebSearchPrice). + for _, item := range summary.ToolSurchargeItems { + surcharge = surcharge.Add(decimal.NewFromFloat(item.Price). + Mul(decimal.NewFromInt(int64(item.Count))). Div(decimal.NewFromInt(1000)). Mul(dGroupRatio). Mul(dQuotaPerUnit)) } - summary.ClaudeWebSearchCallCount = ctx.GetInt("claude_web_search_requests") - if summary.ClaudeWebSearchCallCount > 0 { - summary.ClaudeWebSearchPrice = operation_setting.GetToolPrice("web_search") - surcharge = surcharge.Add(decimal.NewFromFloat(summary.ClaudeWebSearchPrice). - Div(decimal.NewFromInt(1000)). - Mul(dGroupRatio). - Mul(dQuotaPerUnit). - Mul(decimal.NewFromInt(int64(summary.ClaudeWebSearchCallCount)))) - } - - if relayInfo.ResponsesUsageInfo != nil { - if fileSearchTool, exists := relayInfo.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolFileSearch]; exists && fileSearchTool.CallCount > 0 { - summary.FileSearchCallCount = fileSearchTool.CallCount - summary.FileSearchPrice = operation_setting.GetToolPrice("file_search") - surcharge = surcharge.Add(decimal.NewFromFloat(summary.FileSearchPrice). - Mul(decimal.NewFromInt(int64(fileSearchTool.CallCount))). - Div(decimal.NewFromInt(1000)). - Mul(dGroupRatio). - Mul(dQuotaPerUnit)) - } - } - - if ctx.GetBool("image_generation_call") { - summary.ImageGenerationCallPrice = operation_setting.GetGPTImage1PriceOnceCall(ctx.GetString("image_generation_call_quality"), ctx.GetString("image_generation_call_size")) - surcharge = surcharge.Add(decimal.NewFromFloat(summary.ImageGenerationCallPrice). - Mul(dGroupRatio). - Mul(dQuotaPerUnit)) - } - return surcharge } @@ -305,9 +355,9 @@ func calculateTextQuotaSummary(ctx *gin.Context, relayInfo *relaycommon.RelayInf promptQuota := baseTokens.Add(cachedTokensWithRatio).Add(imageTokensWithRatio).Add(cachedCreationTokensWithRatio) completionQuota := dCompletionTokens.Mul(dCompletionRatio) quotaCalculateDecimal := promptQuota.Add(completionQuota).Mul(ratio) - quotaCalculateDecimal = quotaCalculateDecimal.Add(summary.ToolCallSurchargeQuota) quotaCalculateDecimal = quotaCalculateDecimal.Add(audioInputQuota) quotaCalculateDecimal = relayInfo.PriceData.ApplyOtherRatiosToDecimal(quotaCalculateDecimal) + quotaCalculateDecimal = quotaCalculateDecimal.Add(summary.ToolCallSurchargeQuota) if !ratio.IsZero() && quotaCalculateDecimal.LessThanOrEqual(decimal.Zero) { quotaCalculateDecimal = decimal.NewFromInt(1) @@ -317,15 +367,15 @@ func calculateTextQuotaSummary(ctx *gin.Context, relayInfo *relaycommon.RelayInf noteQuotaClamp(relayInfo, clamp) } else { quotaCalculateDecimal := dModelPrice.Mul(dQuotaPerUnit).Mul(dGroupRatio) - quotaCalculateDecimal = quotaCalculateDecimal.Add(summary.ToolCallSurchargeQuota) quotaCalculateDecimal = quotaCalculateDecimal.Add(audioInputQuota) quotaCalculateDecimal = relayInfo.PriceData.ApplyOtherRatiosToDecimal(quotaCalculateDecimal) + quotaCalculateDecimal = quotaCalculateDecimal.Add(summary.ToolCallSurchargeQuota) quota, clamp := common.QuotaFromDecimalChecked(quotaCalculateDecimal) summary.Quota = quota noteQuotaClamp(relayInfo, clamp) } - if summary.TotalTokens == 0 { + if !summary.hasBillableUsage() { summary.Quota = 0 } else if !ratio.IsZero() && summary.Quota == 0 { summary.Quota = 1 @@ -372,23 +422,25 @@ func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, us } } - if summary.WebSearchCallCount > 0 { - extraContent = append(extraContent, fmt.Sprintf("Web Search 调用 %d 次,调用花费 %s", summary.WebSearchCallCount, decimal.NewFromFloat(summary.WebSearchPrice).Mul(decimal.NewFromInt(int64(summary.WebSearchCallCount))).Div(decimal.NewFromInt(1000)).Mul(decimal.NewFromFloat(summary.GroupRatio)).Mul(decimal.NewFromFloat(common.QuotaPerUnit)).String())) - } - if summary.ClaudeWebSearchCallCount > 0 { - extraContent = append(extraContent, fmt.Sprintf("Claude Web Search 调用 %d 次,调用花费 %s", summary.ClaudeWebSearchCallCount, decimal.NewFromFloat(summary.ClaudeWebSearchPrice).Div(decimal.NewFromInt(1000)).Mul(decimal.NewFromFloat(summary.GroupRatio)).Mul(decimal.NewFromFloat(common.QuotaPerUnit)).Mul(decimal.NewFromInt(int64(summary.ClaudeWebSearchCallCount))).String())) - } - if summary.FileSearchCallCount > 0 { - extraContent = append(extraContent, fmt.Sprintf("File Search 调用 %d 次,调用花费 %s", summary.FileSearchCallCount, decimal.NewFromFloat(summary.FileSearchPrice).Mul(decimal.NewFromInt(int64(summary.FileSearchCallCount))).Div(decimal.NewFromInt(1000)).Mul(decimal.NewFromFloat(summary.GroupRatio)).Mul(decimal.NewFromFloat(common.QuotaPerUnit)).String())) + for _, item := range summary.ToolSurchargeItems { + q := decimal.NewFromFloat(item.Price). + Mul(decimal.NewFromInt(int64(item.Count))). + Div(decimal.NewFromInt(1000)). + Mul(decimal.NewFromFloat(summary.GroupRatio)). + Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + extraContent = append(extraContent, fmt.Sprintf( + "%s 调用 %d 次,调用花费 %s", + item.Name, + item.Count, + logger.LogQuota(common.QuotaFromDecimal(q)), + )) } if summary.AudioInputPrice > 0 && summary.AudioTokens > 0 { - extraContent = append(extraContent, fmt.Sprintf("Audio Input 花费 %s", decimal.NewFromFloat(summary.AudioInputPrice).Div(decimal.NewFromInt(1000000)).Mul(decimal.NewFromInt(int64(summary.AudioTokens))).Mul(decimal.NewFromFloat(summary.GroupRatio)).Mul(decimal.NewFromFloat(common.QuotaPerUnit)).String())) - } - if summary.ImageGenerationCallPrice > 0 { - extraContent = append(extraContent, fmt.Sprintf("Image Generation Call 花费 %s", decimal.NewFromFloat(summary.ImageGenerationCallPrice).Mul(decimal.NewFromFloat(summary.GroupRatio)).Mul(decimal.NewFromFloat(common.QuotaPerUnit)).String())) + q := decimal.NewFromFloat(summary.AudioInputPrice).Div(decimal.NewFromInt(1000000)).Mul(decimal.NewFromInt(int64(summary.AudioTokens))).Mul(decimal.NewFromFloat(summary.GroupRatio)).Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + extraContent = append(extraContent, fmt.Sprintf("Audio Input 花费 %s", logger.LogQuota(common.QuotaFromDecimal(q)))) } - if summary.TotalTokens == 0 { + if !summary.hasBillableUsage() { extraContent = append(extraContent, "上游没有返回计费信息,无法扣费(可能是上游超时)") logger.LogError(ctx, fmt.Sprintf("total tokens is 0, cannot consume quota, userId %d, channelId %d, tokenId %d, model %s, pre-consumed quota %d", relayInfo.UserId, relayInfo.ChannelId, relayInfo.TokenId, summary.ModelName, relayInfo.FinalPreConsumedQuota)) } else { @@ -433,29 +485,12 @@ func PostTextConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, us other["image_ratio"] = summary.ImageRatio other["image_output"] = summary.ImageTokens } - if summary.WebSearchCallCount > 0 { - other["web_search"] = true - other["web_search_call_count"] = summary.WebSearchCallCount - other["web_search_price"] = summary.WebSearchPrice - } else if summary.ClaudeWebSearchCallCount > 0 { - other["web_search"] = true - other["web_search_call_count"] = summary.ClaudeWebSearchCallCount - other["web_search_price"] = summary.ClaudeWebSearchPrice - } - if summary.FileSearchCallCount > 0 { - other["file_search"] = true - other["file_search_call_count"] = summary.FileSearchCallCount - other["file_search_price"] = summary.FileSearchPrice - } + appendToolSurchargeLogInfo(other, summary.ToolSurchargeItems) if summary.AudioInputPrice > 0 && summary.AudioTokens > 0 { other["audio_input_seperate_price"] = true other["audio_input_token_count"] = summary.AudioTokens other["audio_input_price"] = summary.AudioInputPrice } - if summary.ImageGenerationCallPrice > 0 { - other["image_generation_call"] = true - other["image_generation_call_price"] = summary.ImageGenerationCallPrice - } if summary.CacheCreationTokens > 0 { other["cache_creation_tokens"] = summary.CacheCreationTokens other["cache_creation_ratio"] = summary.CacheCreationRatio diff --git a/service/text_quota_test.go b/service/text_quota_test.go index 93092125..d9e9a6c6 100644 --- a/service/text_quota_test.go +++ b/service/text_quota_test.go @@ -11,9 +11,13 @@ import ( "github.com/QuantumNous/new-api/dto" "github.com/QuantumNous/new-api/pkg/billingexpr" relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/types" "github.com/gin-gonic/gin" + "github.com/shopspring/decimal" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -546,9 +550,12 @@ func TestComposeTieredTextQuotaKeepsToolCallSurcharges(t *testing.T) { gin.SetMode(gin.TestMode) w := httptest.NewRecorder() ctx, _ := gin.CreateTestContext(w) - ctx.Set("image_generation_call", true) - ctx.Set("image_generation_call_quality", "low") - ctx.Set("image_generation_call_size", "1024x1024") + + // 11 $/1K => 0.011 per completed image output, matching the prior fixed low-tier charge. + operation_setting.SetToolPriceForTest(dto.BuildInToolImageGeneration, 11.0) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest(dto.BuildInToolImageGeneration) + }) relayInfo := &relaycommon.RelayInfo{ OriginModelName: "o1", @@ -559,12 +566,15 @@ func TestComposeTieredTextQuotaKeepsToolCallSurcharges(t *testing.T) { }, ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ - dto.BuildInToolWebSearchPreview: &relaycommon.BuildInToolInfo{ + dto.BuildInToolWebSearchPreview: { CallCount: 1, }, - dto.BuildInToolFileSearch: &relaycommon.BuildInToolInfo{ + dto.BuildInToolFileSearch: { CallCount: 2, }, + dto.BuildInToolImageGeneration: { + CallCount: 1, + }, }, }, TieredBillingSnapshot: &billingexpr.BillingSnapshot{ @@ -740,3 +750,304 @@ func TestCalculateTextQuotaSummaryFixedPriceAppliesImageCountOnceAndAllowsOverri summary = calculateTextQuotaSummary(ctx, relayInfo, usage) require.Equal(t, 120000, summary.Quota) } + +func TestCalculateTextToolCallSurchargeGeneralizedBuiltInTools(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + + operation_setting.SetToolPriceForTest("my_fn", 5.0) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest("my_fn") + }) + + relayInfo := &relaycommon.RelayInfo{ + OriginModelName: "o1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolWebSearchPreview: {CallCount: 2}, + "my_fn": {CallCount: 3}, + "unpriced": {CallCount: 5}, + }, + }, + } + summary := &textQuotaSummary{ + ModelName: "o1", + GroupRatio: 1, + } + + surcharge := calculateTextToolCallSurcharge(ctx, relayInfo, summary) + expected := decimal.NewFromFloat((10.0*2 + 5.0*3) / 1000).Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + assert.True(t, expected.Equal(surcharge), "got %s want %s", surcharge, expected) + require.Len(t, summary.ToolSurchargeItems, 2) + assert.Equal(t, "my_fn", summary.ToolSurchargeItems[0].Name) + assert.Equal(t, 3, summary.ToolSurchargeItems[0].Count) + assert.Equal(t, 5.0, summary.ToolSurchargeItems[0].Price) + assert.Equal(t, dto.BuildInToolWebSearchPreview, summary.ToolSurchargeItems[1].Name) + assert.Equal(t, 2, summary.ToolSurchargeItems[1].Count) + assert.Equal(t, 10.0, summary.ToolSurchargeItems[1].Price) +} + +func TestCalculateTextToolCallSurchargeKeepsSearchPreviewFallbackWithCustomFunctions(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + + operation_setting.SetToolPriceForTest("my_fn", 5) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest("my_fn") + }) + + relayInfo := &relaycommon.RelayInfo{ + RelayMode: relayconstant.RelayModeChatCompletions, + OriginModelName: "gpt-4o-search-preview", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + "my_fn": {CallCount: 1}, + }, + }, + } + summary := &textQuotaSummary{ + ModelName: relayInfo.OriginModelName, + GroupRatio: 1, + } + + surcharge := calculateTextToolCallSurcharge(ctx, relayInfo, summary) + + require.Len(t, summary.ToolSurchargeItems, 2) + assert.Equal(t, "my_fn", summary.ToolSurchargeItems[0].Name) + assert.Equal(t, dto.BuildInToolWebSearchPreview, summary.ToolSurchargeItems[1].Name) + expected := decimal.NewFromFloat((5.0 + 25.0) / 1000). + Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + assert.True(t, expected.Equal(surcharge), "got %s want %s", surcharge, expected) +} + +func TestCalculateTextToolCallSurchargeDoesNotInferSearchForResponses(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + relayInfo := &relaycommon.RelayInfo{ + RelayMode: relayconstant.RelayModeResponses, + OriginModelName: "gpt-4o-search-preview", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{}, + }, + } + summary := &textQuotaSummary{ + ModelName: relayInfo.OriginModelName, + GroupRatio: 1, + } + + surcharge := calculateTextToolCallSurcharge(ctx, relayInfo, summary) + + assert.True(t, surcharge.IsZero()) + assert.Empty(t, summary.ToolSurchargeItems) +} + +func TestCalculateTextToolCallSurchargeMergesSameNameAndPrice(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set("claude_web_search_requests", 3) + + relayInfo := &relaycommon.RelayInfo{ + OriginModelName: "claude-3-7-sonnet", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolWebSearch: {CallCount: 2}, + }, + }, + } + summary := &textQuotaSummary{ModelName: relayInfo.OriginModelName, GroupRatio: 1} + + surcharge := calculateTextToolCallSurcharge(ctx, relayInfo, summary) + + require.Len(t, summary.ToolSurchargeItems, 1) + assert.Equal(t, dto.BuildInToolWebSearch, summary.ToolSurchargeItems[0].Name) + assert.Equal(t, 5, summary.ToolSurchargeItems[0].Count) + assert.Equal(t, 10.0, summary.ToolSurchargeItems[0].Price) + expected := decimal.NewFromFloat(10.0 * 5 / 1000).Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + assert.True(t, expected.Equal(surcharge), "got %s want %s", surcharge, expected) +} + +func TestMergeToolSurchargeItemsSaturatesCountOverflow(t *testing.T) { + items := []ToolSurchargeItem{ + {Name: "custom_fn", Count: math.MaxInt, Price: 5}, + {Name: "custom_fn", Count: 1, Price: 5}, + } + + merged := mergeToolSurchargeItems(items) + + require.Len(t, merged, 1) + assert.Equal(t, math.MaxInt, merged[0].Count) +} + +// A zero-token request (e.g. /v1/alpha/search returns no usage) must still +// bill a tool-call surcharge. Regression for the TotalTokens==0 gate zeroing +// out the surcharge quota. +func TestCalculateTextQuotaSummaryZeroTokensStillBillsToolSurcharge(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + + relayInfo := &relaycommon.RelayInfo{ + OriginModelName: "o1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolWebSearchPreview: {CallCount: 1}, + }, + }, + } + relayInfo.PriceData.GroupRatioInfo.GroupRatio = 1 + + usage := &dto.Usage{} // zero tokens, mirrors alpha search + summary := calculateTextQuotaSummary(ctx, relayInfo, usage) + + require.Equal(t, 0, summary.TotalTokens) + assert.False(t, summary.ToolCallSurchargeQuota.IsZero(), "surcharge should be computed") + assert.Greater(t, summary.Quota, 0, "quota must not be zeroed out for a zero-token web search request") + expected := common.QuotaFromDecimal(summary.ToolCallSurchargeQuota) + assert.Equal(t, expected, summary.Quota) +} + +func TestCalculateTextQuotaSummaryDoesNotApplyRequestMultipliersToToolSurcharge(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + + relayInfo := &relaycommon.RelayInfo{ + OriginModelName: "o1", + PriceData: types.PriceData{ + ModelRatio: 1, + CompletionRatio: 1, + GroupRatioInfo: types.GroupRatioInfo{GroupRatio: 1}, + }, + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolWebSearchPreview: {CallCount: 1}, + }, + }, + } + relayInfo.PriceData.AddOtherRatio("n", 3) + + summary := calculateTextQuotaSummary(ctx, relayInfo, &dto.Usage{}) + + expected := decimal.NewFromFloat(10.0 / 1000).Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + assert.True(t, expected.Equal(summary.ToolCallSurchargeQuota)) + assert.Equal(t, common.QuotaFromDecimal(expected), summary.Quota) +} + +func TestCalculateTextToolCallSurchargeGeminiGoogleSearch(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Set("gemini_google_search_call", true) + + relayInfo := &relaycommon.RelayInfo{OriginModelName: "gemini-2.5-flash"} + summary := &textQuotaSummary{ModelName: "gemini-2.5-flash", GroupRatio: 1} + + surcharge := calculateTextToolCallSurcharge(ctx, relayInfo, summary) + expected := decimal.NewFromFloat(14.0 / 1000).Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + assert.True(t, expected.Equal(surcharge), "got %s want %s", surcharge, expected) + require.Len(t, summary.ToolSurchargeItems, 1) + assert.Equal(t, dto.BuildInToolGoogleSearch, summary.ToolSurchargeItems[0].Name) + assert.Equal(t, 1, summary.ToolSurchargeItems[0].Count) + assert.Equal(t, 14.0, summary.ToolSurchargeItems[0].Price) +} + +func TestCalculateTextToolCallSurchargeImageGenerationDefaultPrice(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest(dto.BuildInToolImageGeneration) + }) + + relayInfo := &relaycommon.RelayInfo{ + OriginModelName: "gpt-5.1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolImageGeneration: {CallCount: 2}, + }, + }, + } + summary := &textQuotaSummary{ModelName: "gpt-5.1", GroupRatio: 1.5} + + surcharge := calculateTextToolCallSurcharge(ctx, relayInfo, summary) + expected := decimal.NewFromFloat(150.0). + Mul(decimal.NewFromInt(2)). + Div(decimal.NewFromInt(1000)). + Mul(decimal.NewFromFloat(1.5)). + Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + assert.True(t, expected.Equal(surcharge), "got %s want %s", surcharge, expected) + require.Len(t, summary.ToolSurchargeItems, 1) + assert.Equal(t, dto.BuildInToolImageGeneration, summary.ToolSurchargeItems[0].Name) + assert.Equal(t, 2, summary.ToolSurchargeItems[0].Count) + assert.Equal(t, 150.0, summary.ToolSurchargeItems[0].Price) +} + +func TestCalculateTextToolCallSurchargeImageGenerationExplicitZeroDisables(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + operation_setting.SetToolPriceForTest(dto.BuildInToolImageGeneration, 0) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest(dto.BuildInToolImageGeneration) + }) + + relayInfo := &relaycommon.RelayInfo{ + OriginModelName: "gpt-5.1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolImageGeneration: {CallCount: 3}, + }, + }, + } + summary := &textQuotaSummary{ModelName: "gpt-5.1", GroupRatio: 1} + + surcharge := calculateTextToolCallSurcharge(ctx, relayInfo, summary) + assert.True(t, surcharge.IsZero()) + assert.Empty(t, summary.ToolSurchargeItems) +} + +func TestCalculateTextQuotaSummaryImageGenerationUsesStructuredSurcharge(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + t.Cleanup(func() { + operation_setting.DeleteToolPriceForTest(dto.BuildInToolImageGeneration) + }) + + relayInfo := &relaycommon.RelayInfo{ + OriginModelName: "gpt-5.1", + ResponsesUsageInfo: &relaycommon.ResponsesUsageInfo{ + BuiltInTools: map[string]*relaycommon.BuildInToolInfo{ + dto.BuildInToolImageGeneration: {CallCount: 1}, + }, + }, + } + relayInfo.PriceData.GroupRatioInfo.GroupRatio = 1 + relayInfo.PriceData.ModelRatio = 1 + relayInfo.PriceData.CompletionRatio = 1 + + usage := &dto.Usage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15} + summary := calculateTextQuotaSummary(ctx, relayInfo, usage) + + require.Len(t, summary.ToolSurchargeItems, 1) + assert.Equal(t, dto.BuildInToolImageGeneration, summary.ToolSurchargeItems[0].Name) + assert.Equal(t, 1, summary.ToolSurchargeItems[0].Count) + assert.Equal(t, 150.0, summary.ToolSurchargeItems[0].Price) + + expectedSurcharge := decimal.NewFromFloat(150.0 / 1000).Mul(decimal.NewFromFloat(common.QuotaPerUnit)) + assert.True(t, expectedSurcharge.Equal(summary.ToolCallSurchargeQuota), + "got %s want %s", summary.ToolCallSurchargeQuota, expectedSurcharge) + assert.Greater(t, summary.Quota, 0) +} + +func TestAppendToolSurchargeLogInfoWritesOnlyStructuredFields(t *testing.T) { + items := []ToolSurchargeItem{ + {Name: dto.BuildInToolWebSearch, Count: 2, Price: 10}, + {Name: dto.BuildInToolImageGeneration, Count: 1, Price: 150}, + } + other := map[string]interface{}{} + + appendToolSurchargeLogInfo(other, items) + + assert.Equal(t, items, other["tool_surcharges"]) + assert.NotContains(t, other, "web_search") + assert.NotContains(t, other, "web_search_call_count") + assert.NotContains(t, other, "web_search_price") + assert.NotContains(t, other, "file_search") + assert.NotContains(t, other, "image_generation_call") + assert.NotContains(t, other, "image_generation_call_price") +} diff --git a/service/tool_billing.go b/service/tool_billing.go deleted file mode 100644 index cadc750c..00000000 --- a/service/tool_billing.go +++ /dev/null @@ -1,86 +0,0 @@ -package service - -import ( - "github.com/QuantumNous/new-api/common" - "github.com/QuantumNous/new-api/setting/operation_setting" -) - -// ToolCallUsage captures all tool call counts from a single request. -type ToolCallUsage struct { - ModelName string - WebSearchCalls int - WebSearchToolName string // "web_search_preview", "web_search", etc. - FileSearchCalls int - ImageGenerationCall bool - ImageGenerationQuality string - ImageGenerationSize string -} - -// ToolCallItem represents a single billed tool usage line. -type ToolCallItem struct { - Name string `json:"name"` - CallCount int `json:"call_count"` - PricePer1K float64 `json:"price_per_1k"` - TotalPrice float64 `json:"total_price"` - Quota int `json:"quota"` -} - -// ToolCallResult holds the aggregated tool call billing for a request. -type ToolCallResult struct { - TotalQuota int `json:"total_quota"` - Items []ToolCallItem `json:"items,omitempty"` -} - -// ComputeToolCallQuota calculates the total quota for all tool calls in a -// request. Tool prices are resolved via GetToolPriceForModel which supports -// model-prefix overrides. groupRatio is applied. -func ComputeToolCallQuota(usage ToolCallUsage, groupRatio float64) ToolCallResult { - var items []ToolCallItem - totalQuota := 0 - - addItem := func(toolName string, count int) { - if count <= 0 { - return - } - pricePer1K := operation_setting.GetToolPriceForModel(toolName, usage.ModelName) - if pricePer1K <= 0 { - return - } - totalPrice := pricePer1K * float64(count) / 1000 - quota := common.QuotaRound(totalPrice * common.QuotaPerUnit * groupRatio) - items = append(items, ToolCallItem{ - Name: toolName, - CallCount: count, - PricePer1K: pricePer1K, - TotalPrice: totalPrice, - Quota: quota, - }) - totalQuota += quota - } - - if usage.WebSearchCalls > 0 && usage.WebSearchToolName != "" { - addItem(usage.WebSearchToolName, usage.WebSearchCalls) - } - - if usage.FileSearchCalls > 0 { - addItem("file_search", usage.FileSearchCalls) - } - - if usage.ImageGenerationCall { - price := operation_setting.GetGPTImage1PriceOnceCall(usage.ImageGenerationQuality, usage.ImageGenerationSize) - quota := common.QuotaRound(price * common.QuotaPerUnit * groupRatio) - items = append(items, ToolCallItem{ - Name: "image_generation", - CallCount: 1, - PricePer1K: price, - TotalPrice: price, - Quota: quota, - }) - totalQuota += quota - } - - return ToolCallResult{ - TotalQuota: totalQuota, - Items: items, - } -} diff --git a/setting/operation_setting/tools.go b/setting/operation_setting/tools.go index 0eb2da0e..07560f89 100644 --- a/setting/operation_setting/tools.go +++ b/setting/operation_setting/tools.go @@ -1,10 +1,14 @@ package operation_setting import ( + "encoding/json" + "fmt" + "math" "sort" "strings" "sync/atomic" + "github.com/QuantumNous/new-api/common" "github.com/QuantumNous/new-api/setting/config" ) @@ -16,39 +20,45 @@ import ( // - "tool_name" → default price for all models // - "tool_name:model_prefix*" → override for models matching the prefix // -// Lookup order: longest prefix match → default → hardcoded fallback → 0 +// Effective index: hardcoded defaults → hardcoded model overrides → valid +// operator values. Lookup uses the longest model prefix before the tool +// default, and a matched numeric zero is terminal. // --------------------------------------------------------------------------- -var defaultToolPrices = map[string]float64{ - "web_search": 10.0, // OpenAI web search (all models) / Claude web search - "web_search_preview": 10.0, // OpenAI web search preview (default: reasoning models) - "file_search": 2.5, // OpenAI file search (Responses API) - "google_search": 14.0, // Gemini Grounding with Google Search -} +const ToolPriceOptionKey = "tool_price_setting.prices" -var defaultToolPriceOverrides = map[string]float64{ - "web_search_preview:gpt-4o*": 25.0, // non-reasoning models - "web_search_preview:gpt-4.1*": 25.0, - "web_search_preview:gpt-4o-mini*": 25.0, - "web_search_preview:gpt-4.1-mini*": 25.0, +const ( + defaultWebSearchToolPrice = 10.0 + defaultWebSearchPreviewToolPrice = 10.0 + defaultFileSearchToolPrice = 2.5 + defaultGoogleSearchToolPrice = 14.0 + defaultImageGenerationToolPrice = 150.0 + defaultSearchPreviewModelPrice = 25.0 +) + +// seedHardcodedToolPrices injects compile-time built-in fallbacks (tool +// defaults and model-prefix overrides) into the destination. The source is +// constants, not a mutable package map or operator configuration. +func seedHardcodedToolPrices(prices map[string]float64) { + prices["web_search"] = defaultWebSearchToolPrice + prices["web_search_preview"] = defaultWebSearchPreviewToolPrice + prices["file_search"] = defaultFileSearchToolPrice + prices["google_search"] = defaultGoogleSearchToolPrice + prices["image_generation"] = defaultImageGenerationToolPrice + prices["web_search_preview:gpt-4o*"] = defaultSearchPreviewModelPrice + prices["web_search_preview:gpt-4.1*"] = defaultSearchPreviewModelPrice + prices["web_search_preview:gpt-4o-mini*"] = defaultSearchPreviewModelPrice + prices["web_search_preview:gpt-4.1-mini*"] = defaultSearchPreviewModelPrice } // ToolPriceSetting is managed by config.GlobalConfig.Register. +// Prices holds operator overrides only; hardcoded fallbacks live in the index. type ToolPriceSetting struct { Prices map[string]float64 `json:"prices"` } var toolPriceSetting = ToolPriceSetting{ - Prices: func() map[string]float64 { - m := make(map[string]float64, len(defaultToolPrices)+len(defaultToolPriceOverrides)) - for k, v := range defaultToolPrices { - m[k] = v - } - for k, v := range defaultToolPriceOverrides { - m[k] = v - } - return m - }(), + Prices: make(map[string]float64), } func init() { @@ -72,17 +82,77 @@ type toolPriceIndex struct { var currentIndex atomic.Pointer[toolPriceIndex] +func isValidToolPrice(price float64) bool { + return price >= 0 && !math.IsNaN(price) && !math.IsInf(price, 0) +} + +func decodeToolPricesJSON(value string, ignoreInvalidEntries bool) (map[string]float64, error) { + rawValue := json.RawMessage(strings.TrimSpace(value)) + if common.GetJsonType(rawValue) != "object" { + return nil, fmt.Errorf("工具价格必须是 JSON 对象") + } + + var rawPrices map[string]json.RawMessage + if err := common.Unmarshal(rawValue, &rawPrices); err != nil { + return nil, fmt.Errorf("解析工具价格失败: %w", err) + } + + prices := make(map[string]float64, len(rawPrices)) + for name, rawPrice := range rawPrices { + var entryErr error + if common.GetJsonType(rawPrice) != "number" { + entryErr = fmt.Errorf("工具价格 %q 必须是非负数字", name) + } else { + var price float64 + if err := common.Unmarshal(rawPrice, &price); err != nil { + entryErr = fmt.Errorf("解析工具价格 %q 失败: %w", name, err) + } else if !isValidToolPrice(price) { + entryErr = fmt.Errorf("工具价格 %q 必须是有限的非负数字", name) + } else { + prices[name] = price + } + } + + if entryErr == nil { + continue + } + if !ignoreInvalidEntries { + return nil, entryErr + } + common.SysError(entryErr.Error()) + } + return prices, nil +} + +// ValidateToolPricesJSON validates an operator-supplied complete price map. +// A numeric zero is valid and intentionally disables the matching rule. +func ValidateToolPricesJSON(value string) error { + _, err := decodeToolPricesJSON(value, false) + return err +} + +// LoadToolPricesFromJSONString replaces the complete operator price map. +// Invalid legacy entries are ignored individually so valid sibling overrides +// survive, while missing built-in keys continue to use hardcoded fallbacks. +func LoadToolPricesFromJSONString(value string) { + prices, err := decodeToolPricesJSON(value, true) + if err != nil { + common.SysError("加载工具价格失败,将使用硬编码兜底: " + err.Error()) + prices = make(map[string]float64) + } + toolPriceSetting.Prices = prices + RebuildToolPriceIndex() +} + // RebuildToolPriceIndex rebuilds the lookup index from the current config. // Called on init and after config updates. Not on the billing hot path. func RebuildToolPriceIndex() { - merged := make(map[string]float64, len(defaultToolPrices)+len(defaultToolPriceOverrides)+len(toolPriceSetting.Prices)) - for k, v := range defaultToolPrices { - merged[k] = v - } - for k, v := range defaultToolPriceOverrides { - merged[k] = v - } + merged := make(map[string]float64, 9+len(toolPriceSetting.Prices)) + seedHardcodedToolPrices(merged) for k, v := range toolPriceSetting.Prices { + if !isValidToolPrice(v) { + continue + } merged[k] = v } @@ -106,6 +176,9 @@ func RebuildToolPriceIndex() { for tool := range idx.prefixes { entries := idx.prefixes[tool] sort.Slice(entries, func(i, j int) bool { + if len(entries[i].prefix) == len(entries[j].prefix) { + return entries[i].prefix < entries[j].prefix + } return len(entries[i].prefix) > len(entries[j].prefix) }) idx.prefixes[tool] = entries @@ -119,10 +192,11 @@ func RebuildToolPriceIndex() { func GetToolPriceForModel(toolName, modelName string) float64 { idx := currentIndex.Load() if idx == nil { - if v, ok := defaultToolPrices[toolName]; ok { - return v + RebuildToolPriceIndex() + idx = currentIndex.Load() + if idx == nil { + return 0 } - return 0 } if entries, ok := idx.prefixes[toolName]; ok && modelName != "" { @@ -144,48 +218,19 @@ func GetToolPrice(toolName string) float64 { return GetToolPriceForModel(toolName, "") } -// --------------------------------------------------------------------------- -// GPT Image 1 per-call pricing (special: depends on quality + size) -// --------------------------------------------------------------------------- - -const ( - GPTImage1Low1024x1024 = 0.011 - GPTImage1Low1024x1536 = 0.016 - GPTImage1Low1536x1024 = 0.016 - GPTImage1Medium1024x1024 = 0.042 - GPTImage1Medium1024x1536 = 0.063 - GPTImage1Medium1536x1024 = 0.063 - GPTImage1High1024x1024 = 0.167 - GPTImage1High1024x1536 = 0.25 - GPTImage1High1536x1024 = 0.25 -) - -func GetGPTImage1PriceOnceCall(quality string, size string) float64 { - prices := map[string]map[string]float64{ - "low": { - "1024x1024": GPTImage1Low1024x1024, - "1024x1536": GPTImage1Low1024x1536, - "1536x1024": GPTImage1Low1536x1024, - }, - "medium": { - "1024x1024": GPTImage1Medium1024x1024, - "1024x1536": GPTImage1Medium1024x1536, - "1536x1024": GPTImage1Medium1536x1024, - }, - "high": { - "1024x1024": GPTImage1High1024x1024, - "1024x1536": GPTImage1High1024x1536, - "1536x1024": GPTImage1High1536x1024, - }, +// SetToolPriceForTest injects a tool price and rebuilds the lookup index. Tests only. +func SetToolPriceForTest(name string, price float64) { + if toolPriceSetting.Prices == nil { + toolPriceSetting.Prices = make(map[string]float64) } + toolPriceSetting.Prices[name] = price + RebuildToolPriceIndex() +} - if qualityMap, exists := prices[quality]; exists { - if price, exists := qualityMap[size]; exists { - return price - } - } - - return GPTImage1High1024x1024 +// DeleteToolPriceForTest removes an injected tool price and rebuilds the index. Tests only. +func DeleteToolPriceForTest(name string) { + delete(toolPriceSetting.Prices, name) + RebuildToolPriceIndex() } // --------------------------------------------------------------------------- @@ -204,15 +249,20 @@ const ( func GetGeminiInputAudioPricePerMillionTokens(modelName string) float64 { if strings.HasPrefix(modelName, "gemini-2.5-flash-preview-native-audio") { return Gemini25FlashNativeAudioInputAudioPrice - } else if strings.HasPrefix(modelName, "gemini-2.5-flash-preview-lite") { + } + if strings.HasPrefix(modelName, "gemini-2.5-flash-preview-lite") { return Gemini25FlashLitePreviewInputAudioPrice - } else if strings.HasPrefix(modelName, "gemini-2.5-flash-preview") { + } + if strings.HasPrefix(modelName, "gemini-2.5-flash-preview") { return Gemini25FlashPreviewInputAudioPrice - } else if strings.HasPrefix(modelName, "gemini-2.5-flash") { + } + if strings.HasPrefix(modelName, "gemini-2.5-flash") { return Gemini25FlashProductionInputAudioPrice - } else if strings.HasPrefix(modelName, "gemini-2.0-flash") { + } + if strings.HasPrefix(modelName, "gemini-2.0-flash") { return Gemini20FlashInputAudioPrice - } else if strings.HasPrefix(modelName, "gemini-robotics-er-1.5") { + } + if strings.HasPrefix(modelName, "gemini-robotics-er-1.5") { return GeminiRoboticsER15InputAudioPrice } return 0 diff --git a/setting/operation_setting/tools_price_test.go b/setting/operation_setting/tools_price_test.go new file mode 100644 index 00000000..027b5192 --- /dev/null +++ b/setting/operation_setting/tools_price_test.go @@ -0,0 +1,155 @@ +package operation_setting + +import ( + "math" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func preserveToolPrices(t *testing.T) { + t.Helper() + original := make(map[string]float64, len(toolPriceSetting.Prices)) + for key, price := range toolPriceSetting.Prices { + original[key] = price + } + t.Cleanup(func() { + toolPriceSetting.Prices = original + RebuildToolPriceIndex() + }) +} + +func TestToolPriceHardcodedFallbacksSurviveMissingOperatorConfig(t *testing.T) { + preserveToolPrices(t) + toolPriceSetting.Prices = map[string]float64{} + RebuildToolPriceIndex() + + expectedDefaults := map[string]float64{ + "web_search": 10, + "web_search_preview": 10, + "file_search": 2.5, + "google_search": 14, + "image_generation": 150, + } + for name, expected := range expectedDefaults { + assert.Equal(t, expected, GetToolPrice(name), name) + } + assert.Equal(t, 25.0, GetToolPriceForModel("web_search_preview", "gpt-4o-2024-11-20")) + assert.Equal(t, 25.0, GetToolPriceForModel("web_search_preview", "gpt-4.1-mini")) +} + +func TestToolPriceOperatorOverridePrecedenceAndExplicitZero(t *testing.T) { + preserveToolPrices(t) + toolPriceSetting.Prices = map[string]float64{ + "image_generation": 0, + "web_search": 12, + "web_search_preview": 0, + "web_search_preview:gpt-4o*": 30, + "web_search_preview:gpt-4o-mini*": 0, + "web_search_preview:custom-model*": 7, + } + RebuildToolPriceIndex() + + assert.Equal(t, 0.0, GetToolPrice("image_generation")) + assert.Equal(t, 12.0, GetToolPrice("web_search")) + assert.Equal(t, 0.0, GetToolPriceForModel("web_search_preview", "o1")) + assert.Equal(t, 30.0, GetToolPriceForModel("web_search_preview", "gpt-4o")) + assert.Equal(t, 0.0, GetToolPriceForModel("web_search_preview", "gpt-4o-mini")) + assert.Equal(t, 25.0, GetToolPriceForModel("web_search_preview", "gpt-4.1")) + assert.Equal(t, 7.0, GetToolPriceForModel("web_search_preview", "custom-model-v2")) + + delete(toolPriceSetting.Prices, "web_search_preview:gpt-4o*") + RebuildToolPriceIndex() + assert.Equal(t, 25.0, GetToolPriceForModel("web_search_preview", "gpt-4o")) + + delete(toolPriceSetting.Prices, "web_search") + RebuildToolPriceIndex() + assert.Equal(t, 10.0, GetToolPrice("web_search")) +} + +func TestToolPriceCustomFunctionHasNoHardcodedFallback(t *testing.T) { + preserveToolPrices(t) + toolPriceSetting.Prices = map[string]float64{} + RebuildToolPriceIndex() + + assert.Equal(t, 0.0, GetToolPrice("lookup_customer")) + + toolPriceSetting.Prices["lookup_customer"] = 5 + RebuildToolPriceIndex() + assert.Equal(t, 5.0, GetToolPrice("lookup_customer")) + + toolPriceSetting.Prices["lookup_customer"] = 0 + RebuildToolPriceIndex() + assert.Equal(t, 0.0, GetToolPrice("lookup_customer")) +} + +func TestValidateToolPricesJSON(t *testing.T) { + valid := []string{ + `{}`, + `{"web_search":0}`, + `{"web_search":10,"custom_fn":2.5}`, + } + for _, value := range valid { + assert.NoError(t, ValidateToolPricesJSON(value), value) + } + + invalid := []string{ + `null`, + `[]`, + `{"web_search":null}`, + `{"web_search":true}`, + `{"web_search":"0"}`, + `{"web_search":-1}`, + `{"web_search":1e999}`, + `{"web_search":`, + } + for _, value := range invalid { + assert.Error(t, ValidateToolPricesJSON(value), value) + } +} + +func TestLoadToolPricesFromJSONStringReplacesMapAndKeepsValidSiblings(t *testing.T) { + preserveToolPrices(t) + + LoadToolPricesFromJSONString(`{ + "web_search": 0, + "custom_fn": 3, + "file_search": null, + "google_search": -1, + "image_generation": "0" + }`) + + require.Len(t, toolPriceSetting.Prices, 2) + assert.Equal(t, 0.0, toolPriceSetting.Prices["web_search"]) + assert.Equal(t, 3.0, toolPriceSetting.Prices["custom_fn"]) + assert.Equal(t, 0.0, GetToolPrice("web_search")) + assert.Equal(t, 3.0, GetToolPrice("custom_fn")) + assert.Equal(t, 2.5, GetToolPrice("file_search")) + assert.Equal(t, 14.0, GetToolPrice("google_search")) + assert.Equal(t, 150.0, GetToolPrice("image_generation")) + + LoadToolPricesFromJSONString(`{"image_generation":0}`) + require.Len(t, toolPriceSetting.Prices, 1) + assert.NotContains(t, toolPriceSetting.Prices, "web_search") + assert.NotContains(t, toolPriceSetting.Prices, "custom_fn") + assert.Equal(t, 10.0, GetToolPrice("web_search")) + assert.Equal(t, 0.0, GetToolPrice("custom_fn")) + assert.Equal(t, 0.0, GetToolPrice("image_generation")) +} + +func TestRebuildToolPriceIndexIgnoresInvalidDirectValues(t *testing.T) { + preserveToolPrices(t) + toolPriceSetting.Prices = map[string]float64{ + "web_search": -1, + "file_search": math.Inf(1), + "image_generation": math.NaN(), + "custom_fn": math.NaN(), + } + RebuildToolPriceIndex() + + assert.Equal(t, 10.0, GetToolPrice("web_search")) + assert.Equal(t, 2.5, GetToolPrice("file_search")) + assert.Equal(t, 150.0, GetToolPrice("image_generation")) + assert.Equal(t, 0.0, GetToolPrice("custom_fn")) +} diff --git a/types/relay_format.go b/types/relay_format.go index 9b4c86f2..86c2de17 100644 --- a/types/relay_format.go +++ b/types/relay_format.go @@ -8,6 +8,7 @@ const ( RelayFormatGemini = "gemini" RelayFormatOpenAIResponses = "openai_responses" RelayFormatOpenAIResponsesCompaction = "openai_responses_compaction" + RelayFormatOpenAIAlphaSearch = "openai_alpha_search" RelayFormatOpenAIAudio = "openai_audio" RelayFormatOpenAIImage = "openai_image" RelayFormatOpenAIRealtime = "openai_realtime" diff --git a/web/src/assets/custom/icon-sub2api.tsx b/web/src/assets/custom/icon-sub2api.tsx new file mode 100644 index 00000000..65a45663 --- /dev/null +++ b/web/src/assets/custom/icon-sub2api.tsx @@ -0,0 +1,62 @@ +/* +Copyright (C) 2023-2026 QuantumNous + +This program is free software: you can redistribute it and/or modify +it under the terms of the GNU Affero General Public License as +published by the Free Software Foundation, either version 3 of the +License, or (at your option) any later version. + +This program is distributed in the hope that it will be useful, +but WITHOUT ANY WARRANTY; without even the implied warranty of +MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +GNU Affero General Public License for more details. + +You should have received a copy of the GNU Affero General Public License +along with this program. If not, see . + +For commercial licensing, please contact support@quantumnous.com +*/ +import { useId, type SVGProps } from 'react' + +type IconSub2apiProps = SVGProps & { + size?: number +} + +export function IconSub2api({ size = 20, ...props }: IconSub2apiProps) { + const gradientId = useId() + + return ( + + + + + + + + + + + + + + ) +} diff --git a/web/src/components/ui/badge.tsx b/web/src/components/ui/badge.tsx index d703d006..96fb780c 100644 --- a/web/src/components/ui/badge.tsx +++ b/web/src/components/ui/badge.tsx @@ -32,6 +32,8 @@ const badgeVariants = cva( 'bg-secondary text-secondary-foreground [a]:hover:bg-secondary/80', destructive: 'bg-destructive/10 text-destructive focus-visible:ring-destructive/20 dark:bg-destructive/20 dark:focus-visible:ring-destructive/40 [a]:hover:bg-destructive/20', + warning: + 'border-warning/40 bg-warning/10 text-warning focus-visible:ring-warning/20', outline: 'border-border text-foreground [a]:hover:bg-muted [a]:hover:text-muted-foreground', ghost: diff --git a/web/src/features/channels/constants.ts b/web/src/features/channels/constants.ts index 556726bc..008494c0 100644 --- a/web/src/features/channels/constants.ts +++ b/web/src/features/channels/constants.ts @@ -77,12 +77,13 @@ export const CHANNEL_TYPES = { 56: 'Replicate', 57: 'ChatGPT Subscription (Codex)', 58: 'Advanced Custom', + 59: 'Sub2API', } as const const CHANNEL_TYPE_DISPLAY_ORDER: number[] = [ 1, 14, 33, 24, 43, 3, 41, 48, 58, 42, 34, 20, 4, 40, 27, 25, 17, 26, 15, 46, - 23, 18, 45, 31, 35, 49, 19, 47, 37, 38, 39, 11, 8, 57, 22, 21, 44, 2, 5, 36, - 50, 51, 52, 53, 54, 55, 56, + 23, 18, 45, 31, 35, 49, 19, 47, 37, 38, 39, 11, 8, 57, 59, 22, 21, 44, 2, 5, + 36, 50, 51, 52, 53, 54, 55, 56, ] export const CHANNEL_TYPE_OPTIONS: { value: number; label: string }[] = (() => { @@ -380,6 +381,7 @@ export const FIELD_DESCRIPTIONS = { export const MODEL_FETCHABLE_TYPES = new Set([ 1, 4, 14, 17, 20, 23, 24, 25, 26, 27, 31, 34, 35, 40, 42, 43, 47, 48, 57, 58, + 59, ]) export const TYPE_TO_KEY_PROMPT: Record = { @@ -391,6 +393,7 @@ export const TYPE_TO_KEY_PROMPT: Record = { 50: 'Format: AccessKey|SecretKey (or just ApiKey if upstream is New API)', 51: 'Format: Access Key ID|Secret Access Key', 57: 'Paste Codex OAuth JSON credential (access_token / refresh_token / account_id)', + 59: 'Enter API key for this channel', } export const CHANNEL_TYPE_WARNINGS: Record = { diff --git a/web/src/features/channels/lib/advanced-custom.ts b/web/src/features/channels/lib/advanced-custom.ts index 302dfdcd..1ab07a50 100644 --- a/web/src/features/channels/lib/advanced-custom.ts +++ b/web/src/features/channels/lib/advanced-custom.ts @@ -107,6 +107,10 @@ export const ADVANCED_CUSTOM_INCOMING_PATH_OPTIONS: AdvancedCustomIncomingPathOp value: '/v1/responses/compact', label: 'OpenAI Responses Compact', }, + { + value: '/v1/alpha/search', + label: 'OpenAI Alpha Search', + }, { value: ADVANCED_CUSTOM_MODEL_LIST_PATH, label: ADVANCED_CUSTOM_MODEL_LIST_LABEL, @@ -834,6 +838,7 @@ function isConverterPathAllowed( converter: AdvancedCustomConverter ): boolean { if (converter === 'none') return true + if (incomingPath === '/v1/alpha/search') return false if (converter === 'anthropic_messages_to_openai_chat_completions') { return incomingPath === '/v1/messages' } diff --git a/web/src/features/channels/lib/channel-type-config.ts b/web/src/features/channels/lib/channel-type-config.ts index 56bf4bce..dec5eb89 100644 --- a/web/src/features/channels/lib/channel-type-config.ts +++ b/web/src/features/channels/lib/channel-type-config.ts @@ -144,6 +144,16 @@ export const CHANNEL_TYPE_CONFIGS: Record = { models: 'Models exposed by this channel', }, }, + 59: { + id: 59, + name: CHANNEL_TYPES[59], + icon: 'Sub2API', + hints: { + baseUrl: 'Sub2API gateway base URL', + key: 'Sub2API API Key', + models: 'Models fetched from upstream /v1/models', + }, + }, } /** diff --git a/web/src/features/channels/lib/channel-utils.ts b/web/src/features/channels/lib/channel-utils.ts index d30c4f7a..98d85251 100644 --- a/web/src/features/channels/lib/channel-utils.ts +++ b/web/src/features/channels/lib/channel-utils.ts @@ -52,6 +52,7 @@ export function getChannelTypeIcon(type: number): string { 7: 'OpenAI', // OhMyGPT 8: 'OpenAI', // Custom 58: 'NewAPI', // Advanced Custom + 59: 'Sub2API', // Sub2API 3: 'Azure', // Azure // Anthropic diff --git a/web/src/features/system-settings/models/__tests__/tool-price-validation.test.tsx b/web/src/features/system-settings/models/__tests__/tool-price-validation.test.tsx new file mode 100644 index 00000000..fc6b8a88 --- /dev/null +++ b/web/src/features/system-settings/models/__tests__/tool-price-validation.test.tsx @@ -0,0 +1,143 @@ +/* +Copyright (C) 2023-2026 QuantumNous + +This program is free software: you can redistribute it and/or modify +it under the terms of the GNU Affero General Public License as +published by the Free Software Foundation, either version 3 of the +License, or (at your option) any later version. + +This program is distributed in the hope that it will be useful, +but WITHOUT ANY WARRANTY; without even the implied warranty of +MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +GNU Affero General Public License for more details. + +You should have received a copy of the GNU Affero General Public License +along with this program. If not, see . + +For commercial licensing, please contact support@quantumnous.com +*/ +import assert from 'node:assert/strict' +import { after, describe, test } from 'node:test' + +import { Window } from 'happy-dom' + +const domWindow = new Window() +const domGlobals = [ + 'window', + 'document', + 'navigator', + 'HTMLElement', + 'HTMLInputElement', + 'SVGElement', + 'Node', + 'Element', + 'Event', + 'CustomEvent', + 'MutationObserver', + 'requestAnimationFrame', + 'cancelAnimationFrame', + 'getComputedStyle', +] as const + +for (const key of domGlobals) { + Object.defineProperty(globalThis, key, { + configurable: true, + value: domWindow[key], + }) +} + +const { act } = await import('react') +const { createRoot } = await import('react-dom/client') +const { QueryClient, QueryClientProvider } = + await import('@tanstack/react-query') +const { createInstance } = await import('i18next') +const { I18nextProvider, initReactI18next } = await import('react-i18next') +const { ToolPriceSettings } = await import('../tool-price-settings') + +const i18n = createInstance() +await i18n.use(initReactI18next).init({ + lng: 'en', + resources: { + en: { + translation: { + 'Price ($/1K calls)': 'Price ($/1K calls)', + 'Please enter a valid number': 'Please enter a valid number', + 'Tool identifier': 'Tool identifier', + }, + }, + }, +}) + +const reactTestGlobals = globalThis as typeof globalThis & { + IS_REACT_ACT_ENVIRONMENT?: boolean +} +reactTestGlobals.IS_REACT_ACT_ENVIRONMENT = true + +function changeInputValue(input: HTMLInputElement, value: string) { + const valueSetter = Object.getOwnPropertyDescriptor( + domWindow.HTMLInputElement.prototype, + 'value' + )?.set + assert.ok(valueSetter) + valueSetter.call(input, value) + input.dispatchEvent( + new domWindow.Event('input', { bubbles: true }) as unknown as Event + ) +} + +describe('tool price validation', () => { + after(() => { + domWindow.close() + }) + + test('blocks an empty price without converting it to an explicit zero', async () => { + const container = document.createElement('div') + document.body.append(container) + const root = createRoot(container) + const queryClient = new QueryClient({ + defaultOptions: { mutations: { retry: false } }, + }) + + await act(async () => { + root.render( + + + + + + ) + }) + + const priceInput = container.querySelector( + 'input[aria-label="Price ($/1K calls): web_search"]' + ) + assert.ok(priceInput) + + await act(async () => { + changeInputValue(priceInput, '') + }) + + assert.equal(priceInput.getAttribute('aria-invalid'), 'true') + assert.equal( + priceInput.closest('[data-slot="field"]')?.querySelector('[role="alert"]') + ?.textContent, + 'Please enter a valid number' + ) + const saveButton = [...container.querySelectorAll('button')].find( + (button) => button.textContent === 'Save tool prices' + ) + assert.ok(saveButton) + assert.equal(saveButton.disabled, true) + + await act(async () => { + changeInputValue(priceInput, '0') + }) + + assert.equal(priceInput.getAttribute('aria-invalid'), 'false') + assert.equal(saveButton.disabled, false) + + await act(async () => root.unmount()) + container.remove() + queryClient.clear() + }) +}) diff --git a/web/src/features/system-settings/models/tool-price-settings.tsx b/web/src/features/system-settings/models/tool-price-settings.tsx index d674a9d7..c06cfe22 100644 --- a/web/src/features/system-settings/models/tool-price-settings.tsx +++ b/web/src/features/system-settings/models/tool-price-settings.tsx @@ -25,6 +25,7 @@ import { StaticDataTable } from '@/components/data-table' import { JsonCodeEditor } from '@/components/json-code-editor' import { Alert, AlertDescription } from '@/components/ui/alert' import { Button } from '@/components/ui/button' +import { Field, FieldError } from '@/components/ui/field' import { Input } from '@/components/ui/input' import { useUpdateOption } from '../hooks/use-update-option' @@ -40,12 +41,20 @@ const DEFAULT_PRICES: Record = { 'web_search_preview:gpt-4.1-mini*': 25.0, file_search: 2.5, google_search: 14.0, + image_generation: 150.0, } type ToolPriceRow = { id: number key: string - price: number + price: string +} + +function parseToolPrice(value: string): number | null { + if (value.trim() === '') return null + const price = Number(value) + if (!Number.isFinite(price) || price < 0) return null + return price } function rowsToObject(rows: ToolPriceRow[]): Record { @@ -53,7 +62,9 @@ function rowsToObject(rows: ToolPriceRow[]): Record { for (const row of rows) { const k = row.key.trim() if (!k) continue - prices[k] = Number(row.price) || 0 + const price = parseToolPrice(row.price) + if (price === null) continue + prices[k] = price } return prices } @@ -62,7 +73,7 @@ function objectToRows(prices: Record): ToolPriceRow[] { return Object.entries(prices).map(([key, price], index) => ({ id: index + 1, key, - price: Number(price) || 0, + price: String(price), })) } @@ -78,7 +89,18 @@ function parseInitialPrices( !Array.isArray(parsed) && Object.keys(parsed as object).length > 0 ) { - return parsed as Record + const validPrices: Record = {} + for (const [key, value] of Object.entries(parsed)) { + if (typeof value === 'number' && Number.isFinite(value) && value >= 0) { + validPrices[key] = value + } + } + // Merge defaults first so newly introduced tools appear for old stored + // configs, while explicit stored values (including 0) still win. + return { + ...DEFAULT_PRICES, + ...validPrices, + } } } catch { // fall through to defaults @@ -111,6 +133,15 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({ }, [defaultValue]) const currentPrices = useMemo(() => rowsToObject(rows), [rows]) + const invalidRowIds = useMemo( + () => + new Set( + rows + .filter((row) => parseToolPrice(row.price) === null) + .map((row) => row.id) + ), + [rows] + ) const syncFromRows = useCallback((nextRows: ToolPriceRow[]) => { setRows(nextRows) @@ -127,7 +158,19 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({ setJsonError(t('JSON must be an object')) return } - const nextRows = objectToRows(parsed as Record) + const prices: Record = {} + for (const [key, value] of Object.entries(parsed)) { + if ( + typeof value !== 'number' || + !Number.isFinite(value) || + value < 0 + ) { + setJsonError(t('Please enter a valid number')) + return + } + prices[key] = value + } + const nextRows = objectToRows(prices) setRows(nextRows) setNextRowId(nextRows.length + 1) setJsonError('') @@ -139,7 +182,7 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({ ) const updateRow = useCallback( - (id: number, field: 'key' | 'price', value: string | number) => { + (id: number, field: 'key' | 'price', value: string) => { syncFromRows( rows.map((r) => (r.id === id ? { ...r, [field]: value } : r)) ) @@ -148,7 +191,7 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({ ) const addRow = useCallback(() => { - const newRow: ToolPriceRow = { id: nextRowId, key: '', price: 0 } + const newRow: ToolPriceRow = { id: nextRowId, key: '', price: '0' } setNextRowId((prev) => prev + 1) syncFromRows([...rows, newRow]) }, [nextRowId, rows, syncFromRows]) @@ -178,6 +221,10 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({ }, [jsonText, t]) const handleSave = useCallback(async () => { + if (invalidRowIds.size > 0) { + toast.error(t('Please enter a valid number')) + return + } if (editMode === 'json' && jsonError) { toast.error(t('Please fix JSON errors before saving')) return @@ -186,7 +233,7 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({ key: OPTION_KEY, value: JSON.stringify(currentPrices), }) - }, [currentPrices, editMode, jsonError, t, updateOption]) + }, [currentPrices, editMode, invalidRowIds.size, jsonError, t, updateOption]) const toggleEditMode = useCallback(() => { setEditMode((prev) => (prev === 'visual' ? 'json' : 'visual')) @@ -276,17 +323,29 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({ id: 'price', header: t('Price ($/1K calls)'), className: 'w-[200px]', - cell: (row) => ( - - updateRow(row.id, 'price', Number(e.target.value) || 0) - } - /> - ), + cell: (row) => { + const isInvalid = invalidRowIds.has(row.id) + return ( + + + updateRow(row.id, 'price', e.target.value) + } + /> + {isInvalid ? ( + + {t('Please enter a valid number')} + + ) : null} + + ) + }, }, { id: 'actions', @@ -322,7 +381,9 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({