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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions controller/option.go
Original file line number Diff line number Diff line change
Expand Up @@ -151,6 +151,15 @@ func UpdateOption(c *gin.Context) {
})
return
}
case "ImageOutputRatio":
err = ratio_setting.UpdateImageOutputRatioByJSONString(option.Value.(string))
if err != nil {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "图片输出倍率设置失败: " + err.Error(),
})
return
}
Comment on lines +154 to +162

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟠 Major

Don’t mutate runtime ratio state during controller validation.

Line 155 applies the update before model.UpdateOption persists it. If DB/update fails later, memory state can diverge from stored state.

🔧 Suggested direction
 case "ImageOutputRatio":
-	err = ratio_setting.UpdateImageOutputRatioByJSONString(option.Value.(string))
+	err = ratio_setting.CheckImageOutputRatioByJSONString(option.Value.(string))
 	if err != nil {
 		c.JSON(http.StatusOK, gin.H{
 			"success": false,
 			"message": "图片输出倍率设置失败: " + err.Error(),
 		})
 		return
 	}

Then keep the actual mutation only in model.UpdateOption -> updateOptionMap.

🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@controller/option.go` around lines 154 - 162, The controller is mutating
runtime ratio state by calling
ratio_setting.UpdateImageOutputRatioByJSONString(option.Value.(string)) before
persistence; instead, change this flow so the controller only validates the
incoming ratio JSON (e.g., parse/validate the payload or call a non-mutating
validator) and do not call the mutating UpdateImageOutputRatioByJSONString here;
let model.UpdateOption (and its updateOptionMap path) be the single place that
applies the runtime mutation after the DB update succeeds. Concretely: remove
the direct mutation call from the controller, replace it with a validation step
(or a new ratio_setting.ValidateImageOutputRatioJSON function) to ensure the
value is well-formed, and ensure model.UpdateOption/updateOptionMap performs the
actual ratio_setting.UpdateImageOutputRatioByJSONString call after successful
persistence.

case "AudioRatio":
err = ratio_setting.UpdateAudioRatioByJSONString(option.Value.(string))
if err != nil {
Expand Down
23 changes: 12 additions & 11 deletions dto/gemini.go
Original file line number Diff line number Diff line change
Expand Up @@ -458,17 +458,18 @@ type GeminiChatResponse struct {
}

type GeminiUsageMetadata struct {
PromptTokenCount int `json:"promptTokenCount"`
ToolUsePromptTokenCount int `json:"toolUsePromptTokenCount"`
CandidatesTokenCount int `json:"candidatesTokenCount"`
TotalTokenCount int `json:"totalTokenCount"`
ThoughtsTokenCount int `json:"thoughtsTokenCount"`
CachedContentTokenCount int `json:"cachedContentTokenCount"`
PromptTokensDetails []GeminiPromptTokensDetails `json:"promptTokensDetails"`
ToolUsePromptTokensDetails []GeminiPromptTokensDetails `json:"toolUsePromptTokensDetails"`
}

type GeminiPromptTokensDetails struct {
PromptTokenCount int `json:"promptTokenCount"`
ToolUsePromptTokenCount int `json:"toolUsePromptTokenCount"`
CandidatesTokenCount int `json:"candidatesTokenCount"`
TotalTokenCount int `json:"totalTokenCount"`
ThoughtsTokenCount int `json:"thoughtsTokenCount"`
CachedContentTokenCount int `json:"cachedContentTokenCount"`
PromptTokensDetails []GeminiTokensDetails `json:"promptTokensDetails"`
ToolUsePromptTokensDetails []GeminiTokensDetails `json:"toolUsePromptTokensDetails"`
CandidatesTokensDetails []GeminiTokensDetails `json:"candidatesTokensDetails"`
}

type GeminiTokensDetails struct {
Modality string `json:"modality"`
TokenCount int `json:"tokenCount"`
}
Expand Down
1 change: 1 addition & 0 deletions dto/openai_response.go
Original file line number Diff line number Diff line change
Expand Up @@ -259,6 +259,7 @@ type InputTokenDetails struct {

type OutputTokenDetails struct {
TextTokens int `json:"text_tokens"`
ImageTokens int `json:"image_tokens"`
AudioTokens int `json:"audio_tokens"`
ReasoningTokens int `json:"reasoning_tokens"`
}
Expand Down
3 changes: 3 additions & 0 deletions model/option.go
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,7 @@ func InitOptionMap() {
common.OptionMap["UserUsableGroups"] = setting.UserUsableGroups2JSONString()
common.OptionMap["CompletionRatio"] = ratio_setting.CompletionRatio2JSONString()
common.OptionMap["ImageRatio"] = ratio_setting.ImageRatio2JSONString()
common.OptionMap["ImageOutputRatio"] = ratio_setting.ImageOutputRatio2JSONString()
common.OptionMap["AudioRatio"] = ratio_setting.AudioRatio2JSONString()
common.OptionMap["AudioCompletionRatio"] = ratio_setting.AudioCompletionRatio2JSONString()
common.OptionMap["TopUpLink"] = common.TopUpLink
Expand Down Expand Up @@ -432,6 +433,8 @@ func updateOptionMap(key string, value string) (err error) {
err = ratio_setting.UpdateCreateCacheRatioByJSONString(value)
case "ImageRatio":
err = ratio_setting.UpdateImageRatioByJSONString(value)
case "ImageOutputRatio":
err = ratio_setting.UpdateImageOutputRatioByJSONString(value)
case "AudioRatio":
err = ratio_setting.UpdateAudioRatioByJSONString(value)
case "AudioCompletionRatio":
Expand Down
6 changes: 3 additions & 3 deletions model/user_oauth_binding.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,9 +10,9 @@ import (
// UserOAuthBinding stores the binding relationship between users and custom OAuth providers
type UserOAuthBinding struct {
Id int `json:"id" gorm:"primaryKey"`
UserId int `json:"user_id" gorm:"not null;uniqueIndex:ux_user_provider"` // User ID - one binding per user per provider
ProviderId int `json:"provider_id" gorm:"not null;uniqueIndex:ux_user_provider;uniqueIndex:ux_provider_userid"` // Custom OAuth provider ID
ProviderUserId string `json:"provider_user_id" gorm:"type:varchar(256);not null;uniqueIndex:ux_provider_userid"` // User ID from OAuth provider - one OAuth account per provider
UserId int `json:"user_id" gorm:"not null;uniqueIndex:ux_user_provider"` // User ID - one binding per user per provider
ProviderId int `json:"provider_id" gorm:"not null;uniqueIndex:ux_user_provider;uniqueIndex:ux_provider_userid"` // Custom OAuth provider ID
ProviderUserId string `json:"provider_user_id" gorm:"type:varchar(256);not null;uniqueIndex:ux_provider_userid"` // User ID from OAuth provider - one OAuth account per provider
CreatedAt time.Time `json:"created_at"`
}

Expand Down
78 changes: 62 additions & 16 deletions relay/channel/gemini/relay-gemini.go
Original file line number Diff line number Diff line change
Expand Up @@ -1047,17 +1047,23 @@ func buildUsageFromGeminiMetadata(metadata dto.GeminiUsageMetadata, fallbackProm
usage.PromptTokensDetails.CachedTokens = metadata.CachedContentTokenCount

for _, detail := range metadata.PromptTokensDetails {
if detail.Modality == "AUDIO" {
switch detail.Modality {
case "AUDIO":
usage.PromptTokensDetails.AudioTokens += detail.TokenCount
} else if detail.Modality == "TEXT" {
case "TEXT":
usage.PromptTokensDetails.TextTokens += detail.TokenCount
case "IMAGE":
usage.PromptTokensDetails.ImageTokens += detail.TokenCount
}
}
for _, detail := range metadata.ToolUsePromptTokensDetails {
if detail.Modality == "AUDIO" {
switch detail.Modality {
case "AUDIO":
usage.PromptTokensDetails.AudioTokens += detail.TokenCount
} else if detail.Modality == "TEXT" {
case "TEXT":
usage.PromptTokensDetails.TextTokens += detail.TokenCount
case "IMAGE":
usage.PromptTokensDetails.ImageTokens += detail.TokenCount
}
}

Expand Down Expand Up @@ -1281,10 +1287,36 @@ func handleFinalStream(c *gin.Context, info *relaycommon.RelayInfo, resp *dto.Ch
return nil
}

func applyGeminiTokensDetailsToPromptUsage(usage *dto.Usage, details []dto.GeminiTokensDetails) {
for _, detail := range details {
switch detail.Modality {
case "AUDIO":
usage.PromptTokensDetails.AudioTokens = detail.TokenCount
case "TEXT":
usage.PromptTokensDetails.TextTokens = detail.TokenCount
case "IMAGE":
usage.PromptTokensDetails.ImageTokens = detail.TokenCount
}
}
}

func applyGeminiTokensDetailsToCompletionUsage(usage *dto.Usage, details []dto.GeminiTokensDetails) {
for _, detail := range details {
switch detail.Modality {
case "AUDIO":
usage.CompletionTokenDetails.AudioTokens = detail.TokenCount
case "TEXT":
usage.CompletionTokenDetails.TextTokens = detail.TokenCount
case "IMAGE":
usage.CompletionTokenDetails.ImageTokens = detail.TokenCount
}
}
}

func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response, callback func(data string, geminiResponse *dto.GeminiChatResponse) bool) (*dto.Usage, *types.NewAPIError) {
var usage = &dto.Usage{}
var imageCount int
responseText := strings.Builder{}
imageCount := 0

helper.StreamScannerHandler(c, resp, info, func(data string) bool {
var geminiResponse dto.GeminiChatResponse
Expand All @@ -1298,31 +1330,44 @@ func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http
common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason))
}

// 统计图片数量
for _, candidate := range geminiResponse.Candidates {
for _, part := range candidate.Content.Parts {
if part.InlineData != nil && part.InlineData.MimeType != "" {
imageCount++
}
if part.Text != "" {
responseText.WriteString(part.Text)
}
if part.InlineData != nil && strings.HasPrefix(part.InlineData.MimeType, "image") {
imageCount++
}
}
}

// 更新使用量统计
if geminiResponse.UsageMetadata.TotalTokenCount != 0 {
mappedUsage := buildUsageFromGeminiMetadata(geminiResponse.UsageMetadata, info.GetEstimatePromptTokens())
mappedUsage := buildUsageFromGeminiMetadata(geminiResponse.UsageMetadata, info.GetEstimatePromptTokens())
applyGeminiTokensDetailsToCompletionUsage(&mappedUsage, geminiResponse.UsageMetadata.CandidatesTokensDetails)
if mappedUsage.TotalTokens != 0 ||
mappedUsage.PromptTokens != 0 ||
mappedUsage.CompletionTokens != 0 ||
mappedUsage.CompletionTokenDetails.ReasoningTokens != 0 ||
mappedUsage.PromptTokensDetails.TextTokens != 0 ||
mappedUsage.PromptTokensDetails.AudioTokens != 0 ||
mappedUsage.PromptTokensDetails.ImageTokens != 0 ||
mappedUsage.PromptTokensDetails.CachedTokens != 0 {
*usage = mappedUsage
}

return callback(data, &geminiResponse)
})

if imageCount != 0 {
if usage.CompletionTokens == 0 {
usage.CompletionTokens = imageCount * 1400
}
if imageCount != 0 && usage.CompletionTokens == 0 {
usage.CompletionTokens = imageCount * 1400
}
if usage.TotalTokens > 0 && usage.PromptTokens > 0 && usage.CompletionTokens <= 0 {
usage.CompletionTokens = usage.TotalTokens - usage.PromptTokens
}
if usage.PromptTokens > 0 && usage.PromptTokensDetails.TextTokens == 0 && usage.PromptTokensDetails.AudioTokens == 0 && usage.PromptTokensDetails.ImageTokens == 0 {
usage.PromptTokensDetails.TextTokens = usage.PromptTokens
}
if usage.TotalTokens == 0 && usage.PromptTokens > 0 && usage.CompletionTokens > 0 {
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
}

if usage.CompletionTokens <= 0 {
Expand Down Expand Up @@ -1478,6 +1523,7 @@ func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.R
fullTextResponse := responseGeminiChat2OpenAI(c, &geminiResponse)
fullTextResponse.Model = info.UpstreamModelName
usage := buildUsageFromGeminiMetadata(geminiResponse.UsageMetadata, info.GetEstimatePromptTokens())
applyGeminiTokensDetailsToCompletionUsage(&usage, geminiResponse.UsageMetadata.CandidatesTokensDetails)

fullTextResponse.Usage = usage

Expand Down
7 changes: 7 additions & 0 deletions relay/channel/openai/relay-openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,10 @@ import (
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/relay/channel/openrouter"
relaycommon "github.com/QuantumNous/new-api/relay/common"
relayconstant "github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/setting/ratio_setting"

"github.com/QuantumNous/new-api/types"

Expand Down Expand Up @@ -589,6 +591,11 @@ func OpenaiHandlerWithUsage(c *gin.Context, info *relaycommon.RelayInfo, resp *h
usageResp.PromptTokensDetails.ImageTokens += usageResp.InputTokensDetails.ImageTokens
usageResp.PromptTokensDetails.TextTokens += usageResp.InputTokensDetails.TextTokens
}
if (info.RelayMode == relayconstant.RelayModeImagesGenerations || info.RelayMode == relayconstant.RelayModeImagesEdits) && usageResp.OutputTokens > 0 {
if _, ok := ratio_setting.GetImageOutputRatio(info.OriginModelName); ok {
usageResp.CompletionTokenDetails.ImageTokens += usageResp.OutputTokens
}
}
applyUsagePostProcessing(info, &usageResp.Usage, responseBody)
return &usageResp.Usage, nil
}
Expand Down
19 changes: 17 additions & 2 deletions relay/compatible_handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -242,6 +242,7 @@ func postConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage
cacheTokens := usage.PromptTokensDetails.CachedTokens
imageTokens := usage.PromptTokensDetails.ImageTokens
audioTokens := usage.PromptTokensDetails.AudioTokens
completionImageTokens := usage.CompletionTokenDetails.ImageTokens
completionTokens := usage.CompletionTokens
cachedCreationTokens := usage.PromptTokensDetails.CachedCreationTokens

Expand All @@ -251,6 +252,7 @@ func postConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage
completionRatio := relayInfo.PriceData.CompletionRatio
cacheRatio := relayInfo.PriceData.CacheRatio
imageRatio := relayInfo.PriceData.ImageRatio
imageOutputRatio := relayInfo.PriceData.ImageOutputRatio
modelRatio := relayInfo.PriceData.ModelRatio
groupRatio := relayInfo.PriceData.GroupRatioInfo.GroupRatio
modelPrice := relayInfo.PriceData.ModelPrice
Expand All @@ -262,10 +264,12 @@ func postConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage
dImageTokens := decimal.NewFromInt(int64(imageTokens))
dAudioTokens := decimal.NewFromInt(int64(audioTokens))
dCompletionTokens := decimal.NewFromInt(int64(completionTokens))
dCompletionImageTokens := decimal.NewFromInt(int64(completionImageTokens))
dCachedCreationTokens := decimal.NewFromInt(int64(cachedCreationTokens))
dCompletionRatio := decimal.NewFromFloat(completionRatio)
dCacheRatio := decimal.NewFromFloat(cacheRatio)
dImageRatio := decimal.NewFromFloat(imageRatio)
dImageOutputRatio := decimal.NewFromFloat(imageOutputRatio)
dModelRatio := decimal.NewFromFloat(modelRatio)
dGroupRatio := decimal.NewFromFloat(groupRatio)
dModelPrice := decimal.NewFromFloat(modelPrice)
Expand Down Expand Up @@ -378,7 +382,14 @@ func postConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage
Add(imageTokensWithRatio).
Add(dCachedCreationTokensWithRatio)

completionQuota := dCompletionTokens.Mul(dCompletionRatio)
baseCompletionTokens := dCompletionTokens
var completionImageTokensWithRatio decimal.Decimal
if !dCompletionImageTokens.IsZero() {
baseCompletionTokens = baseCompletionTokens.Sub(dCompletionImageTokens)
completionImageTokensWithRatio = dCompletionImageTokens.Mul(dImageOutputRatio)
}

completionQuota := baseCompletionTokens.Mul(dCompletionRatio).Add(completionImageTokensWithRatio)

quotaCalculateDecimal = promptQuota.Add(completionQuota).Mul(ratio)

Expand Down Expand Up @@ -451,7 +462,11 @@ func postConsumeQuota(ctx *gin.Context, relayInfo *relaycommon.RelayInfo, usage
if imageTokens != 0 {
other["image"] = true
other["image_ratio"] = imageRatio
other["image_output"] = imageTokens
other["image_input_tokens"] = imageTokens
}
if completionImageTokens != 0 {
other["completion_image_tokens"] = completionImageTokens
other["image_output_ratio"] = imageOutputRatio
}
if cachedCreationTokens != 0 {
other["cache_creation_tokens"] = cachedCreationTokens
Expand Down
3 changes: 3 additions & 0 deletions relay/helper/price.go
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,7 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
var cacheCreationRatio1h float64
var audioRatio float64
var audioCompletionRatio float64
var imageOutputRatio float64
var freeModel bool
if !usePrice {
preConsumedTokens := common.Max(promptTokens, common.PreConsumedQuota)
Expand All @@ -85,6 +86,7 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
// 固定1h和5min缓存写入价格的比例
cacheCreationRatio1h = cacheCreationRatio * claudeCacheCreation1hMultiplier
imageRatio, _ = ratio_setting.GetImageRatio(info.OriginModelName)
imageOutputRatio, _ = ratio_setting.GetImageOutputRatio(info.OriginModelName)
audioRatio = ratio_setting.GetAudioRatio(info.OriginModelName)
audioCompletionRatio = ratio_setting.GetAudioCompletionRatio(info.OriginModelName)
ratio := modelRatio * groupRatioInfo.GroupRatio
Expand Down Expand Up @@ -124,6 +126,7 @@ func ModelPriceHelper(c *gin.Context, info *relaycommon.RelayInfo, promptTokens
UsePrice: usePrice,
CacheRatio: cacheRatio,
ImageRatio: imageRatio,
ImageOutputRatio: imageOutputRatio,
AudioRatio: audioRatio,
AudioCompletionRatio: audioCompletionRatio,
CacheCreationRatio: cacheCreationRatio,
Expand Down
12 changes: 7 additions & 5 deletions service/task_billing_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -125,8 +125,8 @@ func makeTask(userId, channelId, quota, tokenId int, billingSource string, subsc
SubscriptionId: subscriptionId,
TokenId: tokenId,
BillingContext: &model.TaskBillingContext{
ModelPrice: 0.02,
GroupRatio: 1.0,
ModelPrice: 0.02,
GroupRatio: 1.0,
OriginModelName: "test-model",
},
},
Expand Down Expand Up @@ -615,9 +615,11 @@ type mockAdaptor struct {
adjustReturn int
}

func (m *mockAdaptor) Init(_ *relaycommon.RelayInfo) {}
func (m *mockAdaptor) FetchTask(string, string, map[string]any, string) (*http.Response, error) { return nil, nil }
func (m *mockAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) { return nil, nil }
func (m *mockAdaptor) Init(_ *relaycommon.RelayInfo) {}
func (m *mockAdaptor) FetchTask(string, string, map[string]any, string) (*http.Response, error) {
return nil, nil
}
func (m *mockAdaptor) ParseTaskResult([]byte) (*relaycommon.TaskInfo, error) { return nil, nil }
func (m *mockAdaptor) AdjustBillingOnComplete(_ *model.Task, _ *relaycommon.TaskInfo) int {
return m.adjustReturn
}
Expand Down
14 changes: 9 additions & 5 deletions setting/ratio_setting/exposed_cache.go
Original file line number Diff line number Diff line change
Expand Up @@ -42,11 +42,15 @@ func GetExposedData() gin.H {
return cloneGinH(c.data)
}
newData := gin.H{
"model_ratio": GetModelRatioCopy(),
"completion_ratio": GetCompletionRatioCopy(),
"cache_ratio": GetCacheRatioCopy(),
"create_cache_ratio": GetCreateCacheRatioCopy(),
"model_price": GetModelPriceCopy(),
"model_ratio": GetModelRatioCopy(),
"completion_ratio": GetCompletionRatioCopy(),
"cache_ratio": GetCacheRatioCopy(),
"create_cache_ratio": GetCreateCacheRatioCopy(),
"model_price": GetModelPriceCopy(),
"image_ratio": GetImageRatioCopy(),
"image_output_ratio": GetImageOutputRatioCopy(),
"audio_ratio": GetAudioRatioCopy(),
"audio_completion_ratio": GetAudioCompletionRatioCopy(),
}
exposedData.Store(&exposedCache{
data: newData,
Expand Down
Loading