feat: configurable tool pricing, Sub2API channel, and alpha search billing
Add admin-configurable tool-call prices with cross-provider surcharge settlement, Sub2API channel support, /v1/alpha/search relay, and usage-log surcharge UI.
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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"},
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -37,5 +37,6 @@ const (
|
||||
APITypeReplicate
|
||||
APITypeCodex
|
||||
APITypeAdvancedCustom
|
||||
APITypeSub2API
|
||||
APITypeDummy // this one is only for count, do not add any channel after this
|
||||
)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+9
-4
@@ -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 {
|
||||
|
||||
+8
-37
@@ -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 (
|
||||
|
||||
+20
-2
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,6 @@
|
||||
package sub2api
|
||||
|
||||
const ChannelName = "sub2api"
|
||||
|
||||
// ModelList is empty because models are fetched dynamically from upstream /v1/models.
|
||||
var ModelList = []string{}
|
||||
@@ -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
|
||||
//}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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") {
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
+146
-111
@@ -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
|
||||
|
||||
+316
-5
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
@@ -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"
|
||||
|
||||
@@ -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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
For commercial licensing, please contact support@quantumnous.com
|
||||
*/
|
||||
import { useId, type SVGProps } from 'react'
|
||||
|
||||
type IconSub2apiProps = SVGProps<SVGSVGElement> & {
|
||||
size?: number
|
||||
}
|
||||
|
||||
export function IconSub2api({ size = 20, ...props }: IconSub2apiProps) {
|
||||
const gradientId = useId()
|
||||
|
||||
return (
|
||||
<svg
|
||||
xmlns='http://www.w3.org/2000/svg'
|
||||
viewBox='0 0 24 24'
|
||||
width={size}
|
||||
height={size}
|
||||
{...props}
|
||||
>
|
||||
<defs>
|
||||
<linearGradient
|
||||
id={gradientId}
|
||||
x1='4'
|
||||
y1='4'
|
||||
x2='20'
|
||||
y2='20'
|
||||
gradientUnits='userSpaceOnUse'
|
||||
>
|
||||
<stop stopColor='#67EDB1' />
|
||||
<stop offset='.48' stopColor='#2FD3E1' />
|
||||
<stop offset='1' stopColor='#2E68EA' />
|
||||
</linearGradient>
|
||||
</defs>
|
||||
<g
|
||||
fill='none'
|
||||
stroke={`url(#${gradientId})`}
|
||||
strokeLinecap='round'
|
||||
strokeLinejoin='round'
|
||||
strokeWidth='2.7'
|
||||
>
|
||||
<path d='m19.25 7.65-2.55-3.2a1.33 1.33 0 0 0-1.03-.5H8.58c-.34 0-.67.13-.91.37L4.15 7.65c-.93.88-.6 1.52.65 2.3l9.55 5.97' />
|
||||
<path d='m4.75 16.35 2.55 3.2c.25.31.63.5 1.03.5h7.09c.34 0 .67-.13.91-.37l3.52-3.33c.93-.88.6-1.52-.65-2.3L9.65 8.08' />
|
||||
</g>
|
||||
</svg>
|
||||
)
|
||||
}
|
||||
@@ -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:
|
||||
|
||||
@@ -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<number, string> = {
|
||||
@@ -391,6 +393,7 @@ export const TYPE_TO_KEY_PROMPT: Record<number, string> = {
|
||||
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<number, string> = {
|
||||
|
||||
@@ -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'
|
||||
}
|
||||
|
||||
@@ -144,6 +144,16 @@ export const CHANNEL_TYPE_CONFIGS: Record<number, ChannelTypeConfig> = {
|
||||
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',
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
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(
|
||||
<QueryClientProvider client={queryClient}>
|
||||
<I18nextProvider i18n={i18n}>
|
||||
<ToolPriceSettings defaultValue='{"web_search":10}' />
|
||||
</I18nextProvider>
|
||||
</QueryClientProvider>
|
||||
)
|
||||
})
|
||||
|
||||
const priceInput = container.querySelector<HTMLInputElement>(
|
||||
'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()
|
||||
})
|
||||
})
|
||||
@@ -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<string, number> = {
|
||||
'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<string, number> {
|
||||
@@ -53,7 +62,9 @@ function rowsToObject(rows: ToolPriceRow[]): Record<string, number> {
|
||||
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<string, number>): 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<string, number>
|
||||
const validPrices: Record<string, number> = {}
|
||||
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<string, number>)
|
||||
const prices: Record<string, number> = {}
|
||||
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) => (
|
||||
<Input
|
||||
type='number'
|
||||
min={0}
|
||||
step={0.5}
|
||||
value={row.price}
|
||||
onChange={(e) =>
|
||||
updateRow(row.id, 'price', Number(e.target.value) || 0)
|
||||
}
|
||||
/>
|
||||
),
|
||||
cell: (row) => {
|
||||
const isInvalid = invalidRowIds.has(row.id)
|
||||
return (
|
||||
<Field data-invalid={isInvalid}>
|
||||
<Input
|
||||
type='number'
|
||||
min={0}
|
||||
step={0.5}
|
||||
value={row.price}
|
||||
aria-invalid={isInvalid}
|
||||
aria-label={`${t('Price ($/1K calls)')}: ${row.key || t('Tool identifier')}`}
|
||||
onChange={(e) =>
|
||||
updateRow(row.id, 'price', e.target.value)
|
||||
}
|
||||
/>
|
||||
{isInvalid ? (
|
||||
<FieldError>
|
||||
{t('Please enter a valid number')}
|
||||
</FieldError>
|
||||
) : null}
|
||||
</Field>
|
||||
)
|
||||
},
|
||||
},
|
||||
{
|
||||
id: 'actions',
|
||||
@@ -322,7 +381,9 @@ export const ToolPriceSettings = memo(function ToolPriceSettings({
|
||||
<Button
|
||||
onClick={handleSave}
|
||||
disabled={
|
||||
updateOption.isPending || (editMode === 'json' && !!jsonError)
|
||||
updateOption.isPending ||
|
||||
invalidRowIds.size > 0 ||
|
||||
(editMode === 'json' && !!jsonError)
|
||||
}
|
||||
>
|
||||
{t('Save tool prices')}
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
/*
|
||||
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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
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'
|
||||
import type React from 'react'
|
||||
|
||||
const domWindow = new Window()
|
||||
const domGlobals = [
|
||||
'window',
|
||||
'document',
|
||||
'navigator',
|
||||
'HTMLElement',
|
||||
'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 { createInstance } = await import('i18next')
|
||||
const { I18nextProvider, initReactI18next } = await import('react-i18next')
|
||||
|
||||
const i18n = createInstance()
|
||||
await i18n.use(initReactI18next).init({
|
||||
lng: 'en',
|
||||
resources: {
|
||||
en: {
|
||||
translation: {
|
||||
Subscription: 'Subscription',
|
||||
'Deducted by subscription': 'Deducted by subscription',
|
||||
'Includes tool-call surcharge': 'Includes tool-call surcharge',
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
const { LogCostDisplay } = await import('../log-cost-display')
|
||||
const { formatLogQuota } = await import('@/lib/format')
|
||||
const reactTestGlobals = globalThis as typeof globalThis & {
|
||||
IS_REACT_ACT_ENVIRONMENT?: boolean
|
||||
}
|
||||
reactTestGlobals.IS_REACT_ACT_ENVIRONMENT = true
|
||||
|
||||
type RenderedCost = {
|
||||
container: HTMLDivElement
|
||||
root: ReturnType<typeof createRoot>
|
||||
}
|
||||
|
||||
async function renderCost(
|
||||
props: React.ComponentProps<typeof LogCostDisplay>
|
||||
): Promise<RenderedCost> {
|
||||
const container = document.createElement('div')
|
||||
document.body.append(container)
|
||||
const root = createRoot(container)
|
||||
|
||||
await act(async () => {
|
||||
root.render(
|
||||
<I18nextProvider i18n={i18n}>
|
||||
<LogCostDisplay {...props} />
|
||||
</I18nextProvider>
|
||||
)
|
||||
})
|
||||
|
||||
return { container, root }
|
||||
}
|
||||
|
||||
async function unmountCost(rendered: RenderedCost) {
|
||||
await act(async () => rendered.root.unmount())
|
||||
rendered.container.remove()
|
||||
}
|
||||
|
||||
function normalizedText(value: string | null): string {
|
||||
return (value ?? '').replaceAll(/\s/g, '')
|
||||
}
|
||||
|
||||
describe('log cost display', () => {
|
||||
after(() => {
|
||||
domWindow.close()
|
||||
})
|
||||
|
||||
test('keeps the regular cost visible and adds an accessible surcharge marker', async () => {
|
||||
const rendered = await renderCost({
|
||||
quota: 12500,
|
||||
other: {
|
||||
tool_surcharges: [{ name: 'lookup_customer', count: 1, price: 5 }],
|
||||
},
|
||||
})
|
||||
|
||||
assert.equal(
|
||||
normalizedText(rendered.container.textContent).includes(
|
||||
normalizedText(formatLogQuota(12500))
|
||||
),
|
||||
true
|
||||
)
|
||||
const marker = rendered.container.querySelector(
|
||||
'[data-tool-surcharge-indicator="true"]'
|
||||
)
|
||||
assert.ok(marker)
|
||||
assert.equal(
|
||||
marker.getAttribute('aria-label'),
|
||||
'Includes tool-call surcharge'
|
||||
)
|
||||
assert.equal(marker.getAttribute('tabindex'), '0')
|
||||
|
||||
await unmountCost(rendered)
|
||||
})
|
||||
|
||||
test('preserves the subscription badge and adds the same legacy surcharge marker', async () => {
|
||||
const rendered = await renderCost({
|
||||
quota: 5000,
|
||||
other: {
|
||||
billing_source: 'subscription',
|
||||
web_search: true,
|
||||
web_search_call_count: 1,
|
||||
web_search_price: 10,
|
||||
},
|
||||
})
|
||||
|
||||
assert.equal(rendered.container.textContent?.includes('Subscription'), true)
|
||||
assert.ok(
|
||||
rendered.container.querySelector('[data-tool-surcharge-indicator="true"]')
|
||||
)
|
||||
|
||||
await unmountCost(rendered)
|
||||
})
|
||||
})
|
||||
@@ -58,6 +58,7 @@ import {
|
||||
} from '../../lib/utils'
|
||||
import type { LogOtherData } from '../../types'
|
||||
import { DetailsDialog } from '../dialogs/details-dialog'
|
||||
import { LogCostDisplay } from '../log-cost-display'
|
||||
import { ModelBadge } from '../model-badge'
|
||||
import { TimingMetricsCell, StreamTpsCell } from '../timing-metrics-cell'
|
||||
import { useUsageLogsContext } from '../usage-logs-provider'
|
||||
@@ -93,12 +94,6 @@ function getGroupRatio(other: LogOtherData | null): number | null {
|
||||
return null
|
||||
}
|
||||
|
||||
function splitQuotaDisplay(value: string): { prefix: string; amount: string } {
|
||||
const match = value.match(/^([^0-9+\-.,\s]+)(.+)$/)
|
||||
if (!match) return { prefix: '', amount: value }
|
||||
return { prefix: match[1], amount: match[2] }
|
||||
}
|
||||
|
||||
function buildDetailSegments(
|
||||
log: UsageLog,
|
||||
other: LogOtherData | null,
|
||||
@@ -703,46 +698,7 @@ export function useCommonLogsColumns(isAdmin: boolean): ColumnDef<UsageLog>[] {
|
||||
|
||||
const quota = row.getValue('quota') as number
|
||||
const other = parseLogOther(log.other)
|
||||
const isSubscription = other?.billing_source === 'subscription'
|
||||
|
||||
if (isSubscription) {
|
||||
return (
|
||||
<TooltipProvider>
|
||||
<Tooltip>
|
||||
<TooltipTrigger
|
||||
render={
|
||||
<StatusBadge
|
||||
label={t('Subscription')}
|
||||
variant='success'
|
||||
size='sm'
|
||||
copyable={false}
|
||||
className='cursor-help'
|
||||
/>
|
||||
}
|
||||
/>
|
||||
<TooltipContent>
|
||||
<span>
|
||||
{t('Deducted by subscription')}: {formatLogQuota(quota)}
|
||||
</span>
|
||||
</TooltipContent>
|
||||
</Tooltip>
|
||||
</TooltipProvider>
|
||||
)
|
||||
}
|
||||
|
||||
const quotaStr = formatLogQuota(quota)
|
||||
const quotaDisplay = splitQuotaDisplay(quotaStr)
|
||||
|
||||
return (
|
||||
<div className='flex flex-col gap-0.5'>
|
||||
<span className='border-border/80 bg-muted/60 inline-flex h-6 w-fit items-center rounded-md border px-2 [font-family:var(--font-body)] text-sm leading-none font-semibold tabular-nums'>
|
||||
{quotaDisplay.prefix && (
|
||||
<span className='mr-1'>{quotaDisplay.prefix}</span>
|
||||
)}
|
||||
<span>{quotaDisplay.amount}</span>
|
||||
</span>
|
||||
</div>
|
||||
)
|
||||
return <LogCostDisplay quota={quota} other={other} />
|
||||
},
|
||||
},
|
||||
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
/*
|
||||
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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
For commercial licensing, please contact support@quantumnous.com
|
||||
*/
|
||||
import { Wrench01Icon } from '@hugeicons/core-free-icons'
|
||||
import { HugeiconsIcon } from '@hugeicons/react'
|
||||
import { useTranslation } from 'react-i18next'
|
||||
|
||||
import { StatusBadge } from '@/components/status-badge'
|
||||
import { Badge } from '@/components/ui/badge'
|
||||
import {
|
||||
Tooltip,
|
||||
TooltipContent,
|
||||
TooltipProvider,
|
||||
TooltipTrigger,
|
||||
} from '@/components/ui/tooltip'
|
||||
import { formatLogQuota } from '@/lib/format'
|
||||
|
||||
import { hasToolSurcharge } from '../lib/format'
|
||||
import type { LogOtherData } from '../types'
|
||||
|
||||
interface LogCostDisplayProps {
|
||||
quota: number
|
||||
other: LogOtherData | null
|
||||
}
|
||||
|
||||
function splitQuotaDisplay(value: string): { prefix: string; amount: string } {
|
||||
const match = value.match(/^([^0-9+\-.,\s]+)(.+)$/)
|
||||
if (!match) return { prefix: '', amount: value }
|
||||
return { prefix: match[1], amount: match[2] }
|
||||
}
|
||||
|
||||
function ToolSurchargeMarker() {
|
||||
const { t } = useTranslation()
|
||||
const label = t('Includes tool-call surcharge')
|
||||
|
||||
return (
|
||||
<Tooltip>
|
||||
<TooltipTrigger
|
||||
render={
|
||||
<Badge
|
||||
variant='warning'
|
||||
className='h-5 min-w-5 cursor-help gap-0 rounded-full px-1'
|
||||
role='img'
|
||||
aria-label={label}
|
||||
tabIndex={0}
|
||||
data-tool-surcharge-indicator='true'
|
||||
>
|
||||
<HugeiconsIcon
|
||||
icon={Wrench01Icon}
|
||||
strokeWidth={2}
|
||||
aria-hidden='true'
|
||||
/>
|
||||
<span
|
||||
className='text-[9px] leading-none font-bold'
|
||||
aria-hidden='true'
|
||||
>
|
||||
+
|
||||
</span>
|
||||
</Badge>
|
||||
}
|
||||
/>
|
||||
<TooltipContent>{label}</TooltipContent>
|
||||
</Tooltip>
|
||||
)
|
||||
}
|
||||
|
||||
function QuotaBadge(props: { quota: number }) {
|
||||
const quotaDisplay = splitQuotaDisplay(formatLogQuota(props.quota))
|
||||
|
||||
return (
|
||||
<span className='border-border/80 bg-muted/60 inline-flex h-6 w-fit items-center rounded-md border px-2 [font-family:var(--font-body)] text-sm leading-none font-semibold tabular-nums'>
|
||||
{quotaDisplay.prefix ? (
|
||||
<span className='mr-1'>{quotaDisplay.prefix}</span>
|
||||
) : null}
|
||||
<span>{quotaDisplay.amount}</span>
|
||||
</span>
|
||||
)
|
||||
}
|
||||
|
||||
function SubscriptionBadge(props: { quota: number }) {
|
||||
const { t } = useTranslation()
|
||||
|
||||
return (
|
||||
<Tooltip>
|
||||
<TooltipTrigger
|
||||
render={
|
||||
<StatusBadge
|
||||
label={t('Subscription')}
|
||||
variant='success'
|
||||
size='sm'
|
||||
copyable={false}
|
||||
className='cursor-help'
|
||||
/>
|
||||
}
|
||||
/>
|
||||
<TooltipContent>
|
||||
<span>
|
||||
{t('Deducted by subscription')}: {formatLogQuota(props.quota)}
|
||||
</span>
|
||||
</TooltipContent>
|
||||
</Tooltip>
|
||||
)
|
||||
}
|
||||
|
||||
export function LogCostDisplay(props: LogCostDisplayProps) {
|
||||
const isSubscription = props.other?.billing_source === 'subscription'
|
||||
const showToolSurcharge = hasToolSurcharge(props.other)
|
||||
|
||||
if (!isSubscription && !showToolSurcharge) {
|
||||
return (
|
||||
<div className='flex flex-col gap-0.5'>
|
||||
<QuotaBadge quota={props.quota} />
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
return (
|
||||
<TooltipProvider>
|
||||
<div className='inline-flex items-center gap-1'>
|
||||
{isSubscription ? (
|
||||
<SubscriptionBadge quota={props.quota} />
|
||||
) : (
|
||||
<QuotaBadge quota={props.quota} />
|
||||
)}
|
||||
{showToolSurcharge ? <ToolSurchargeMarker /> : null}
|
||||
</div>
|
||||
</TooltipProvider>
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
/*
|
||||
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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
For commercial licensing, please contact support@quantumnous.com
|
||||
*/
|
||||
import assert from 'node:assert/strict'
|
||||
import { describe, test } from 'node:test'
|
||||
|
||||
import type { LogOtherData } from '../../types'
|
||||
import { hasToolSurcharge } from '../format'
|
||||
|
||||
describe('tool surcharge detection', () => {
|
||||
test('shows the marker for a charged structured tool surcharge', () => {
|
||||
assert.equal(
|
||||
hasToolSurcharge({
|
||||
tool_surcharges: [{ name: 'lookup_customer', count: 2, price: 5 }],
|
||||
}),
|
||||
true
|
||||
)
|
||||
})
|
||||
|
||||
const legacyCases: Array<{
|
||||
name: string
|
||||
other: LogOtherData
|
||||
}> = [
|
||||
{
|
||||
name: 'Web Search',
|
||||
other: {
|
||||
web_search: true,
|
||||
web_search_call_count: 1,
|
||||
web_search_price: 10,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: 'File Search',
|
||||
other: {
|
||||
file_search: true,
|
||||
file_search_call_count: 2,
|
||||
file_search_price: 2.5,
|
||||
},
|
||||
},
|
||||
{
|
||||
name: 'Image Generation',
|
||||
other: {
|
||||
image_generation_call: true,
|
||||
image_generation_call_price: 0.04,
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
for (const scenario of legacyCases) {
|
||||
test(`keeps the marker visible for legacy ${scenario.name} charges`, () => {
|
||||
assert.equal(hasToolSurcharge(scenario.other), true)
|
||||
})
|
||||
}
|
||||
|
||||
test('hides the marker when surcharge entries are empty or not chargeable', () => {
|
||||
const invalidCases: Array<LogOtherData | null> = [
|
||||
null,
|
||||
{},
|
||||
{ tool_surcharges: [] },
|
||||
{
|
||||
tool_surcharges: [{ name: 'lookup_customer', count: 0, price: 5 }],
|
||||
},
|
||||
{
|
||||
tool_surcharges: [{ name: 'lookup_customer', count: 1, price: 0 }],
|
||||
},
|
||||
{
|
||||
tool_surcharges: [{ name: ' ', count: 1, price: 5 }],
|
||||
},
|
||||
{
|
||||
web_search: true,
|
||||
web_search_call_count: 1,
|
||||
web_search_price: 0,
|
||||
},
|
||||
{
|
||||
image_generation_call: false,
|
||||
image_generation_call_price: 0.04,
|
||||
},
|
||||
]
|
||||
|
||||
for (const other of invalidCases) {
|
||||
assert.equal(hasToolSurcharge(other), false)
|
||||
}
|
||||
})
|
||||
})
|
||||
@@ -92,6 +92,67 @@ export function isViolationFeeLog(other: LogOtherData | null): boolean {
|
||||
)
|
||||
}
|
||||
|
||||
function isPositiveFiniteNumber(value: unknown): value is number {
|
||||
return typeof value === 'number' && Number.isFinite(value) && value > 0
|
||||
}
|
||||
|
||||
function hasLegacySearchSurcharge(
|
||||
enabled: boolean | undefined,
|
||||
count: number | undefined,
|
||||
price: number | undefined
|
||||
): boolean {
|
||||
return (
|
||||
enabled === true &&
|
||||
isPositiveFiniteNumber(count) &&
|
||||
isPositiveFiniteNumber(price)
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Check whether a consume log includes an actual tool-call surcharge.
|
||||
* Structured surcharge items cover current logs, while the legacy fields keep
|
||||
* historical Web Search, File Search, and Image Generation logs visible.
|
||||
*/
|
||||
export function hasToolSurcharge(other: LogOtherData | null): boolean {
|
||||
if (!other) return false
|
||||
|
||||
const hasStructuredSurcharge =
|
||||
Array.isArray(other.tool_surcharges) &&
|
||||
other.tool_surcharges.some(
|
||||
(item) =>
|
||||
typeof item?.name === 'string' &&
|
||||
item.name.trim() !== '' &&
|
||||
isPositiveFiniteNumber(item.count) &&
|
||||
isPositiveFiniteNumber(item.price)
|
||||
)
|
||||
if (hasStructuredSurcharge) return true
|
||||
|
||||
if (
|
||||
hasLegacySearchSurcharge(
|
||||
other.web_search,
|
||||
other.web_search_call_count,
|
||||
other.web_search_price
|
||||
)
|
||||
) {
|
||||
return true
|
||||
}
|
||||
|
||||
if (
|
||||
hasLegacySearchSurcharge(
|
||||
other.file_search,
|
||||
other.file_search_call_count,
|
||||
other.file_search_price
|
||||
)
|
||||
) {
|
||||
return true
|
||||
}
|
||||
|
||||
return (
|
||||
other.image_generation_call === true &&
|
||||
isPositiveFiniteNumber(other.image_generation_call_price)
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* Parse the 'other' field from JSON string to object
|
||||
*/
|
||||
|
||||
@@ -106,6 +106,12 @@ export const USAGE_BILLING_PATH = {
|
||||
export type UsageBillingPath =
|
||||
(typeof USAGE_BILLING_PATH)[keyof typeof USAGE_BILLING_PATH]
|
||||
|
||||
export interface ToolSurchargeItem {
|
||||
name: string
|
||||
count: number
|
||||
price: number
|
||||
}
|
||||
|
||||
export interface LogOtherData {
|
||||
admin_info?: {
|
||||
is_multi_key?: boolean
|
||||
@@ -197,11 +203,13 @@ export interface LogOtherData {
|
||||
file_search?: boolean
|
||||
file_search_call_count?: number
|
||||
file_search_price?: number
|
||||
tool_surcharges?: ToolSurchargeItem[]
|
||||
audio_input_seperate_price?: boolean
|
||||
audio_input_token_count?: number
|
||||
audio_input_price?: number
|
||||
image_generation_call?: boolean
|
||||
image_generation_call_price?: number
|
||||
image_generation_call_count?: number
|
||||
is_system_prompt_overwritten?: boolean
|
||||
po?: string[]
|
||||
billing_source?: string
|
||||
|
||||
@@ -2291,6 +2291,7 @@
|
||||
"Include Model": "Include Model",
|
||||
"Include Rule Name": "Include Rule Name",
|
||||
"Includes request rules": "Includes request rules",
|
||||
"Includes tool-call surcharge": "Includes tool-call surcharge",
|
||||
"Including failed requests, 0 = unlimited": "Including failed requests, 0 = unlimited",
|
||||
"Incoming path": "Incoming path",
|
||||
"Incoming path is required": "Incoming path is required",
|
||||
|
||||
@@ -2291,6 +2291,7 @@
|
||||
"Include Model": "Inclure le modèle",
|
||||
"Include Rule Name": "Inclure le nom de la règle",
|
||||
"Includes request rules": "Inclut des règles de requête",
|
||||
"Includes tool-call surcharge": "Inclut un supplément pour appel d’outil",
|
||||
"Including failed requests, 0 = unlimited": "Y compris les requêtes échouées, 0 = illimité",
|
||||
"Incoming path": "Chemin entrant",
|
||||
"Incoming path is required": "Le chemin entrant est requis",
|
||||
|
||||
@@ -2291,6 +2291,7 @@
|
||||
"Include Model": "モデルを含む",
|
||||
"Include Rule Name": "ルール名を含む",
|
||||
"Includes request rules": "リクエストルールを含む",
|
||||
"Includes tool-call surcharge": "ツール呼び出しの追加料金を含む",
|
||||
"Including failed requests, 0 = unlimited": "失敗したリクエストを含む、0 = 無制限",
|
||||
"Incoming path": "受信パス",
|
||||
"Incoming path is required": "受信パスは必須です",
|
||||
|
||||
@@ -2291,6 +2291,7 @@
|
||||
"Include Model": "Включить модель",
|
||||
"Include Rule Name": "Включить имя правила",
|
||||
"Includes request rules": "Включает правила запросов",
|
||||
"Includes tool-call surcharge": "Включает доплату за вызов инструмента",
|
||||
"Including failed requests, 0 = unlimited": "Включая неудачные запросы, 0 = без ограничений",
|
||||
"Incoming path": "Входящий путь",
|
||||
"Incoming path is required": "Входящий путь обязателен",
|
||||
|
||||
@@ -2291,6 +2291,7 @@
|
||||
"Include Model": "Bao gồm mô hình",
|
||||
"Include Rule Name": "Bao gồm tên quy tắc",
|
||||
"Includes request rules": "Bao gồm quy tắc yêu cầu",
|
||||
"Includes tool-call surcharge": "Bao gồm phụ phí gọi công cụ",
|
||||
"Including failed requests, 0 = unlimited": "Bao gồm các yêu cầu thất bại, 0 = không giới hạn",
|
||||
"Incoming path": "Path đầu vào",
|
||||
"Incoming path is required": "Bắt buộc path đầu vào",
|
||||
|
||||
@@ -2291,6 +2291,7 @@
|
||||
"Include Model": "包含模型",
|
||||
"Include Rule Name": "包含規則名",
|
||||
"Includes request rules": "包含請求規則",
|
||||
"Includes tool-call surcharge": "包含工具呼叫附加費",
|
||||
"Including failed requests, 0 = unlimited": "包括失敗的請求,0 = 無限制",
|
||||
"Incoming path": "入口路徑",
|
||||
"Incoming path is required": "入口路徑不能為空",
|
||||
|
||||
@@ -2291,6 +2291,7 @@
|
||||
"Include Model": "包含模型",
|
||||
"Include Rule Name": "包含规则名",
|
||||
"Includes request rules": "包含请求规则",
|
||||
"Includes tool-call surcharge": "包含工具调用附加费",
|
||||
"Including failed requests, 0 = unlimited": "包括失败的请求,0 = 无限制",
|
||||
"Incoming path": "入口路径",
|
||||
"Incoming path is required": "入口路径不能为空",
|
||||
|
||||
@@ -26,6 +26,13 @@ For commercial licensing, please contact support@quantumnous.com
|
||||
* - Size parameter: getLobeIcon("OpenAI", 20)
|
||||
*/
|
||||
import * as LobeIcons from '@lobehub/icons'
|
||||
import type React from 'react'
|
||||
|
||||
import { IconSub2api } from '@/assets/custom/icon-sub2api'
|
||||
|
||||
const CUSTOM_ICONS: Record<string, React.ComponentType<{ size?: number }>> = {
|
||||
Sub2API: IconSub2api,
|
||||
}
|
||||
|
||||
/**
|
||||
* Parse a property value from string to appropriate type
|
||||
@@ -102,6 +109,10 @@ export function getLobeIcon(
|
||||
// Parse component path and chained properties
|
||||
const segments = trimmedName.split('.')
|
||||
const baseKey = segments[0]
|
||||
const CustomIcon = CUSTOM_ICONS[baseKey]
|
||||
if (CustomIcon) {
|
||||
return <CustomIcon size={size} />
|
||||
}
|
||||
const BaseIcon = (LobeIcons as Record<string, unknown>)[baseKey] as
|
||||
| Record<string, unknown>
|
||||
| undefined
|
||||
|
||||
Reference in New Issue
Block a user