fix(gemini): 原生生图按上游实际回吐的图片张数计费

/v1beta/models/{model}:generateContent 与 Anthropic→Gemini 兼容路径的
ImageCount 只由 isImageGenerationModel(originalModel) 决定,而该白名单是
按 Google 官方模型名精确/前缀匹配写死的(antigravity_image_test.go 里
明确断言 my-gemini-3-pro-image-test 这类自定义名返回 false)。

GeminiMessagesCompatService 服务的却主要是 API Key + 自定义模型映射的账号:
客户端请求名和 GetMappedModel 后的上游名都可能是站长自取的别名,白名单必然
判不出来 → ImageCount=0 → calculateRecordUsageCost 里 `if result.ImageCount > 0`
的按次计费分支整条不触发 → 生图请求全部记 $0(issue #5358)。

改为优先按上游响应里真实的 inlineData 图片 part 计数:
- 新增请求级计数器,挂在 gin.Context 上,与既有的
  upstreamResponseModelObserver 同一批调用点取解包后的响应体;
- 取「单个 payload 内的最大值」而非累加:Gemini 兼容上游的 SSE 分片可能是
  累积式的(computeGeminiTextDelta 正是为此存在),逐 chunk 累加会把同一张图
  重复计费。max 保证累积式流与非流式都得到真实张数,增量式多图流最差退化到 1,
  与改动前同值,不构成回退;
- 每次 Forward 开头重置计数器,避免 failover 复用同一个 gin.Context 时
  把失败账号已回吐的图叠加到成功账号账单上;
- 响应里数不出图时(fileData 引用式回图等)退回原有模型名启发式,并额外认
  映射后的上游模型名,与 shouldSkipCodexPlanGatedImageModelCooldown 同时取
  requestedModel / modelKey 的口径一致。

inlineData / inline_data 两种字段风格都认(官方 SDK 与部分中转回 snake_case),
只统计带 base64 数据且 MIME 为图片的 part。

Fixes #5358
This commit is contained in:
li
2026-08-08 21:36:18 +08:00
parent cc67b1aca1
commit b6eb6c1efa
3 changed files with 346 additions and 8 deletions
@@ -0,0 +1,132 @@
package service
import (
"strings"
"github.com/gin-gonic/gin"
"github.com/tidwall/gjson"
)
// geminiImageOutputCounterKey 是请求级内联图片计数器挂在 gin.Context 上的键。
const geminiImageOutputCounterKey = "gemini_image_output_counter"
// geminiImageOutputCounter 记录一次转发里 Gemini 上游真正回吐的内联图片数量。
//
// 取「单个 payload 内的最大值」而不是累加:Gemini 兼容上游的 SSE 分片可能是
// 累积式的(同一份内容在后续 chunk 里重复整段回来——本文件同目录的
// computeGeminiTextDelta 就是为此存在的),逐 chunk 累加会把同一张图算很多次。
// 计费上宁可少算不可多算,所以用 max 兜底:
// - 非流式:整份响应体只观测一次,max 即真实张数;
// - 累积式流:最后一个 chunk 含全部图片,max 仍是真实张数;
// - 增量式流且多图分散在不同 chunk:会低估到 1,与改动前的模型名启发式同值,
// 不构成回退。
type geminiImageOutputCounter struct {
count int
}
// beginGeminiImageOutputObservation 在每次 Forward 开头重置计数器。
// failover 会拿同一个 gin.Context 重跑转发,不重置就会把上一个账号的图数带进来。
func beginGeminiImageOutputObservation(c *gin.Context) *geminiImageOutputCounter {
if c == nil {
return nil
}
counter := &geminiImageOutputCounter{}
c.Set(geminiImageOutputCounterKey, counter)
return counter
}
func geminiImageOutputCounterFromContext(c *gin.Context) *geminiImageOutputCounter {
if c == nil {
return nil
}
value, ok := c.Get(geminiImageOutputCounterKey)
if !ok {
return nil
}
counter, _ := value.(*geminiImageOutputCounter)
return counter
}
// observeGeminiImageOutputs 观测一段上游响应(整份或单个 chunk)里的内联图片。
// 调用点与 upstreamResponseModelObserver.ObserveGemini 一一对应——那里拿得到
// 解包后的上游响应体,这里需要的是同一份字节。
func observeGeminiImageOutputs(c *gin.Context, payload []byte) {
counter := geminiImageOutputCounterFromContext(c)
if counter == nil {
return
}
if count := countGeminiInlineImageOutputs(payload); count > counter.count {
counter.count = count
}
}
func observedGeminiImageOutputs(c *gin.Context) int {
counter := geminiImageOutputCounterFromContext(c)
if counter == nil {
return 0
}
return counter.count
}
// resolveGeminiImageCount 决定本次请求按几张图计费。
//
// 优先用上游真正返回的内联图片数:走 GeminiMessagesCompatService 的账号多是
// API Key + 自定义模型映射,客户端请求名和上游模型名都可能是站长自取的别名
// issue #5358 里的 nana-banana-2),isImageGenerationModel 的白名单必然判不出,
// 于是 ImageCount=0calculateRecordUsageCost 整条按次计费分支不触发,
// 生图请求全部记 $0。
//
// 只有响应里数不出图时(例如上游用 fileData 引用而非 inlineData 回图,或聚合
// 函数丢掉了图片 part)才退回既有的模型名启发式,保证老行为不回退;这里额外
// 也认映射后的上游模型名,与 shouldSkipCodexPlanGatedImageModelCooldown 对
// requestedModel / modelKey 双取的口径一致。
func resolveGeminiImageCount(c *gin.Context, originalModel, mappedModel string) int {
if observed := observedGeminiImageOutputs(c); observed > 0 {
return observed
}
if isImageGenerationModel(originalModel) || isImageGenerationModel(mappedModel) {
return 1
}
return 0
}
// countGeminiInlineImageOutputs 统计一段 Gemini 响应 JSON 里的内联图片 part。
// Gemini REST 回 camelCase 的 inlineData,官方 SDK 与部分中转会回 snake_case
// 的 inline_data,两种都要认。
func countGeminiInlineImageOutputs(payload []byte) int {
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return 0
}
count := 0
gjson.GetBytes(payload, "candidates").ForEach(func(_, candidate gjson.Result) bool {
candidate.Get("content.parts").ForEach(func(_, part gjson.Result) bool {
if geminiPartIsInlineImage(part) {
count++
}
return true
})
return true
})
return count
}
func geminiPartIsInlineImage(part gjson.Result) bool {
inline := part.Get("inlineData")
if !inline.Exists() {
inline = part.Get("inline_data")
}
if !inline.Exists() {
return false
}
mimeType := inline.Get("mimeType")
if !mimeType.Exists() {
mimeType = inline.Get("mime_type")
}
if !isGeminiInlineImageMIMEType(strings.ToLower(strings.TrimSpace(mimeType.String()))) {
return false
}
// 只认真的带上了 base64 数据的 part,空壳 part 不计费。
return strings.TrimSpace(inline.Get("data").String()) != ""
}
@@ -0,0 +1,204 @@
//go:build unit
package service
import (
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
const geminiTestPNG = "iVBORw0KGgoAAAANSUhEUg=="
func newGeminiImageTestContext(t *testing.T) *gin.Context {
t.Helper()
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest(http.MethodPost,
"/v1beta/models/nana-banana-2:generateContent", strings.NewReader("{}"))
return c
}
func geminiImageResponse(parts string) string {
return `{"candidates":[{"content":{"role":"model","parts":[` + parts + `]},"finishReason":"STOP"}]}`
}
func TestCountGeminiInlineImageOutputs(t *testing.T) {
cases := []struct {
name string
payload string
want int
}{
{
name: "camelCase inlineData",
payload: geminiImageResponse(`{"inlineData":{"mimeType":"image/png","data":"` + geminiTestPNG + `"}}`),
want: 1,
},
{
// 官方 SDK 与部分中转把字段回成 snake_case。
name: "snake_case inline_data",
payload: geminiImageResponse(`{"inline_data":{"mime_type":"image/png","data":"` + geminiTestPNG + `"}}`),
want: 1,
},
{
name: "text and image mixed",
payload: geminiImageResponse(`{"text":"here you go"},` +
`{"inlineData":{"mimeType":"image/jpeg","data":"` + geminiTestPNG + `"}}`),
want: 1,
},
{
name: "multiple images",
payload: geminiImageResponse(`{"inlineData":{"mimeType":"image/png","data":"` + geminiTestPNG + `"}},` +
`{"inlineData":{"mimeType":"image/webp","data":"` + geminiTestPNG + `"}}`),
want: 2,
},
{
name: "uppercase mime type",
payload: geminiImageResponse(`{"inlineData":{"mimeType":"IMAGE/PNG","data":"` + geminiTestPNG + `"}}`),
want: 1,
},
{
name: "text only",
payload: geminiImageResponse(`{"text":"no image here"}`),
want: 0,
},
{
// 非图片的内联附件(例如音频)不能按图片计费。
name: "non image mime type",
payload: geminiImageResponse(`{"inlineData":{"mimeType":"audio/mpeg","data":"` + geminiTestPNG + `"}}`),
want: 0,
},
{
name: "empty data is not billable",
payload: geminiImageResponse(`{"inlineData":{"mimeType":"image/png","data":""}}`),
want: 0,
},
{name: "empty payload", payload: "", want: 0},
{name: "invalid json", payload: "not-json", want: 0},
{name: "error response", payload: `{"error":{"code":429,"message":"quota"}}`, want: 0},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
require.Equal(t, tc.want, countGeminiInlineImageOutputs([]byte(tc.payload)))
})
}
}
// 累积式 SSE 会把同一张图在后续 chunk 里整段重发,逐 chunk 累加会重复计费。
// 计数器取单个 payload 内的最大值,正是为了挡住这一点。
func TestObserveGeminiImageOutputs_CumulativeChunksDoNotDoubleCount(t *testing.T) {
c := newGeminiImageTestContext(t)
beginGeminiImageOutputObservation(c)
oneImage := geminiImageResponse(`{"inlineData":{"mimeType":"image/png","data":"` + geminiTestPNG + `"}}`)
for range 4 {
observeGeminiImageOutputs(c, []byte(oneImage))
}
require.Equal(t, 1, observedGeminiImageOutputs(c))
}
func TestObserveGeminiImageOutputs_KeepsLargestChunk(t *testing.T) {
c := newGeminiImageTestContext(t)
beginGeminiImageOutputObservation(c)
observeGeminiImageOutputs(c, []byte(geminiImageResponse(`{"text":"working"}`)))
observeGeminiImageOutputs(c, []byte(geminiImageResponse(
`{"inlineData":{"mimeType":"image/png","data":"`+geminiTestPNG+`"}},`+
`{"inlineData":{"mimeType":"image/png","data":"`+geminiTestPNG+`"}}`)))
// 收尾 chunk 只带 usageMetadata,不能把已数到的张数抹掉。
observeGeminiImageOutputs(c, []byte(`{"usageMetadata":{"promptTokenCount":9}}`))
require.Equal(t, 2, observedGeminiImageOutputs(c))
}
// failover 会拿同一个 gin.Context 重跑 Forward,计数器必须按次重置,
// 否则失败账号已经回吐的图会被叠加到成功账号的账单上。
func TestBeginGeminiImageOutputObservation_ResetsPerForward(t *testing.T) {
c := newGeminiImageTestContext(t)
oneImage := []byte(geminiImageResponse(`{"inlineData":{"mimeType":"image/png","data":"` + geminiTestPNG + `"}}`))
beginGeminiImageOutputObservation(c)
observeGeminiImageOutputs(c, oneImage)
require.Equal(t, 1, observedGeminiImageOutputs(c))
beginGeminiImageOutputObservation(c)
require.Equal(t, 0, observedGeminiImageOutputs(c))
observeGeminiImageOutputs(c, oneImage)
require.Equal(t, 1, observedGeminiImageOutputs(c))
}
// issue #5358:自定义模型名(客户端名与上游映射名都不在白名单里)走 Gemini 原生
// generateContent 生图,改动前 ImageCount 恒为 0calculateRecordUsageCost 的按次
// 计费分支整条不触发,四次生图全部记 $0。
func TestResolveGeminiImageCount(t *testing.T) {
oneImage := []byte(geminiImageResponse(`{"inlineData":{"mimeType":"image/png","data":"` + geminiTestPNG + `"}}`))
textOnly := []byte(geminiImageResponse(`{"text":"hello"}`))
t.Run("custom model name bills by observed images", func(t *testing.T) {
c := newGeminiImageTestContext(t)
beginGeminiImageOutputObservation(c)
observeGeminiImageOutputs(c, oneImage)
require.False(t, isImageGenerationModel("nana-banana-2"), "前置条件:白名单判不出自定义名")
require.Equal(t, 1, resolveGeminiImageCount(c, "nana-banana-2", "nana-banana-2"))
})
t.Run("falls back to requested model name", func(t *testing.T) {
c := newGeminiImageTestContext(t)
beginGeminiImageOutputObservation(c)
observeGeminiImageOutputs(c, textOnly)
require.Equal(t, 1, resolveGeminiImageCount(c, "gemini-3-pro-image-preview", "gemini-3-pro-image-preview"))
})
t.Run("falls back to mapped upstream model name", func(t *testing.T) {
c := newGeminiImageTestContext(t)
beginGeminiImageOutputObservation(c)
observeGeminiImageOutputs(c, textOnly)
require.Equal(t, 1, resolveGeminiImageCount(c, "my-image-alias", "gemini-2.5-flash-image"))
})
t.Run("text model stays unbilled", func(t *testing.T) {
c := newGeminiImageTestContext(t)
beginGeminiImageOutputObservation(c)
observeGeminiImageOutputs(c, textOnly)
require.Equal(t, 0, resolveGeminiImageCount(c, "gemini-2.5-pro", "gemini-2.5-pro"))
})
t.Run("no counter on context degrades to name heuristic", func(t *testing.T) {
c := newGeminiImageTestContext(t)
require.Equal(t, 0, resolveGeminiImageCount(c, "nana-banana-2", "nana-banana-2"))
require.Equal(t, 1, resolveGeminiImageCount(c, "gemini-3-pro-image", "gemini-3-pro-image"))
})
}
// 端到端守住接线:/v1beta/models/{model}:generateContent 的非流式响应体
// 必须真的喂进计数器,否则上面的单测全绿而线上依然记 $0。
func TestHandleNativeNonStreamingResponse_FeedsImageCounter(t *testing.T) {
c := newGeminiImageTestContext(t)
beginGeminiImageOutputObservation(c)
body := geminiImageResponse(`{"inlineData":{"mimeType":"image/png","data":"` + geminiTestPNG + `"}}`)
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(strings.NewReader(body)),
}
svc := &GeminiMessagesCompatService{}
usage, err := svc.handleNativeNonStreamingResponse(c, resp, false)
require.NoError(t, err)
require.NotNil(t, usage)
require.Equal(t, 1, observedGeminiImageOutputs(c))
require.Equal(t, 1, resolveGeminiImageCount(c, "nana-banana-2", "nana-banana-2"))
}
@@ -582,6 +582,7 @@ func (s *GeminiMessagesCompatService) SelectAccountForAIStudioEndpoints(ctx cont
func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Context, account *Account, body []byte) (*ForwardResult, error) {
beginUpstreamResponseModelObservation(c)
beginGeminiImageOutputObservation(c)
startTime := time.Now()
var req struct {
@@ -1074,6 +1075,7 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
}
collectedBytes, _ := json.Marshal(collected)
upstreamResponseModelObserverFromContext(c).ObserveGemini(collectedBytes)
observeGeminiImageOutputs(c, collectedBytes)
claudeResp, usageObj2 := convertGeminiToClaudeMessage(collected, originalModel, collectedBytes, false)
c.JSON(http.StatusOK, claudeResp)
usage = usageObj2
@@ -1089,12 +1091,9 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
}
// 图片生成计费
imageCount := 0
imageInputSize := s.extractImageInputSize(body)
imageSize := normalizeOpenAIImageSizeTier(imageInputSize)
if isImageGenerationModel(originalModel) {
imageCount = 1
}
imageCount := resolveGeminiImageCount(c, originalModel, mappedModel)
return &ForwardResult{
RequestID: requestID,
@@ -1122,6 +1121,7 @@ func isGeminiSignatureRelatedError(respBody []byte) bool {
func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.Context, account *Account, originalModel string, action string, stream bool, body []byte) (*ForwardResult, error) {
beginUpstreamResponseModelObservation(c)
beginGeminiImageOutputObservation(c)
startTime := time.Now()
if strings.TrimSpace(originalModel) == "" {
@@ -1609,6 +1609,7 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
}
b, _ := json.Marshal(collected)
upstreamResponseModelObserverFromContext(c).ObserveGemini(b)
observeGeminiImageOutputs(c, b)
c.Data(http.StatusOK, "application/json", b)
usage = usageObj
} else {
@@ -1625,12 +1626,9 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
}
// 图片生成计费
imageCount := 0
imageInputSize := s.extractImageInputSize(body)
imageSize := normalizeOpenAIImageSizeTier(imageInputSize)
if isImageGenerationModel(originalModel) {
imageCount = 1
}
imageCount := resolveGeminiImageCount(c, originalModel, mappedModel)
return &ForwardResult{
RequestID: requestID,
@@ -2012,6 +2010,7 @@ func (s *GeminiMessagesCompatService) handleNonStreamingResponse(c *gin.Context,
observer = beginUpstreamResponseModelObservation(c)
}
observer.ObserveGemini(unwrappedBody)
observeGeminiImageOutputs(c, unwrappedBody)
var geminiResp map[string]any
if err := json.Unmarshal(unwrappedBody, &geminiResp); err != nil {
@@ -2100,6 +2099,7 @@ func (s *GeminiMessagesCompatService) handleStreamingResponse(c *gin.Context, re
observer = beginUpstreamResponseModelObservation(c)
}
observer.ObserveGemini(unwrappedBytes)
observeGeminiImageOutputs(c, unwrappedBytes)
var geminiResp map[string]any
if err := json.Unmarshal(unwrappedBytes, &geminiResp); err != nil {
@@ -2606,6 +2606,7 @@ func (s *GeminiMessagesCompatService) handleNativeNonStreamingResponse(c *gin.Co
observer = beginUpstreamResponseModelObservation(c)
}
observer.ObserveGemini(respBody)
observeGeminiImageOutputs(c, respBody)
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
@@ -2689,6 +2690,7 @@ func (s *GeminiMessagesCompatService) handleNativeStreamingResponse(c *gin.Conte
usage = u
}
observer.ObserveGemini(rawBytes)
observeGeminiImageOutputs(c, rawBytes)
if firstTokenMs == nil {
ms := int(time.Since(startTime).Milliseconds())