fix(openai): use image tool usage for OAuth billing

This commit is contained in:
wucm667
2026-07-16 08:48:43 +08:00
parent eb2b8632de
commit d22f4d9b5b
2 changed files with 243 additions and 8 deletions
@@ -1072,6 +1072,165 @@ func (s *OpenAIGatewayService) tryWriteOpenAIImagesStreamEvent(
return true
}
func (s *OpenAIGatewayService) parseOpenAIImagesSSEUsageBytes(data []byte, usage *OpenAIUsage) {
s.parseSSEUsageBytes(data, usage)
if usage == nil || !gjson.ValidBytes(data) || gjson.GetBytes(data, "type").String() != "response.completed" {
return
}
if toolUsage, ok := openAIImagesToolUsageFromGJSON(gjson.GetBytes(data, "response.tool_usage.image_gen")); ok {
*usage = toolUsage
}
}
func openAIImagesToolUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) {
if !value.Exists() || !value.IsObject() {
return OpenAIUsage{}, false
}
inputTokens, inputOK := boundedJSONNonNegativeInt(value.Get("input_tokens"))
outputTokens, outputOK := boundedJSONNonNegativeInt(value.Get("output_tokens"))
imageOutputTokens, imageOutputOK := boundedJSONNonNegativeInt(value.Get("output_tokens_details.image_tokens"))
if !inputOK || !outputOK || !imageOutputOK {
return OpenAIUsage{}, false
}
return OpenAIUsage{
InputTokens: inputTokens,
OutputTokens: outputTokens,
ImageOutputTokens: imageOutputTokens,
}, true
}
// boundedJSONNonNegativeInt parses integral JSON exponent notation without
// invoking an arbitrary-precision parser on an upstream-controlled exponent.
func boundedJSONNonNegativeInt(value gjson.Result) (int, bool) {
if !value.Exists() || value.Type != gjson.Number {
return 0, false
}
raw := value.Raw
if len(raw) == 0 || len(raw) > 64 || raw[0] == '-' {
return 0, false
}
mantissaEnd := len(raw)
for i, c := range raw {
if c != 'e' && c != 'E' {
continue
}
mantissaEnd = i
break
}
digits := raw[:mantissaEnd]
fractionDigits := 0
digitCount := 0
dotSeen := false
mantissaIsZero := true
for _, c := range digits {
switch {
case c == '.' && !dotSeen:
dotSeen = true
case c >= '0' && c <= '9':
digitCount++
mantissaIsZero = mantissaIsZero && c == '0'
if dotSeen {
fractionDigits++
}
default:
return 0, false
}
}
exponent := 0
if mantissaEnd < len(raw) {
exponentRaw := raw[mantissaEnd+1:]
negative := false
if len(exponentRaw) > 0 && (exponentRaw[0] == '+' || exponentRaw[0] == '-') {
negative = exponentRaw[0] == '-'
exponentRaw = exponentRaw[1:]
}
if len(exponentRaw) == 0 {
return 0, false
}
for len(exponentRaw) > 1 && exponentRaw[0] == '0' {
exponentRaw = exponentRaw[1:]
}
for _, digit := range exponentRaw {
if digit < '0' || digit > '9' {
return 0, false
}
}
if mantissaIsZero {
return 0, true
}
if len(exponentRaw) > 3 {
return 0, false
}
for _, digit := range exponentRaw {
exponent = exponent*10 + int(digit-'0')
}
if exponent > 100 {
return 0, false
}
if negative {
exponent = -exponent
}
}
trailingZeros := exponent - fractionDigits
scaleReduction := 0
if trailingZeros < 0 {
scaleReduction = -trailingZeros
remaining := scaleReduction
allZeros := true
for i := len(digits) - 1; i >= 0; i-- {
if digits[i] == '.' {
continue
}
if digits[i] != '0' {
allZeros = false
if remaining > 0 {
return 0, false
}
}
if remaining > 0 {
remaining--
}
}
if remaining > 0 {
if allZeros {
return 0, true
}
return 0, false
}
}
maxInt := int(^uint(0) >> 1)
parsed := 0
digitsToAccumulate := digitCount - scaleReduction
for _, c := range digits {
if c == '.' {
continue
}
if digitsToAccumulate <= 0 {
break
}
if parsed > (maxInt-int(c-'0'))/10 {
return 0, false
}
parsed = parsed*10 + int(c-'0')
digitsToAccumulate--
}
if trailingZeros < 0 {
return parsed, true
}
for ; trailingZeros > 0; trailingZeros-- {
if parsed > maxInt/10 {
return 0, false
}
parsed *= 10
}
return parsed, true
}
func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse(
resp *http.Response,
c *gin.Context,
@@ -1085,7 +1244,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse(
var usage OpenAIUsage
forEachOpenAISSEDataPayload(string(body), func(data []byte) {
s.parseSSEUsageBytes(data, &usage)
s.parseOpenAIImagesSSEUsageBytes(data, &usage)
})
results, createdAt, usageRaw, firstMeta, _, err := collectOpenAIImagesFromResponsesBody(body)
if err != nil {
@@ -1192,7 +1351,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
ms := int(time.Since(startTime).Milliseconds())
firstTokenMs = &ms
}
s.parseSSEUsageBytes(dataBytes, &usage)
s.parseOpenAIImagesSSEUsageBytes(dataBytes, &usage)
if !gjson.ValidBytes(dataBytes) {
return
}
+82 -6
View File
@@ -631,7 +631,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthPassesNAndReturnsAllImages(t *te
"X-Request-Id": []string{"req_img_123"},
},
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000000,\"usage\":{\"input_tokens\":11,\"output_tokens\":22,\"input_tokens_details\":{\"cached_tokens\":3},\"output_tokens_details\":{\"image_tokens\":7}},\"tool_usage\":{\"image_gen\":{\"images\":3}},\"output\":[{\"type\":\"image_generation_call\",\"result\":\"aW1hZ2UtMQ==\",\"revised_prompt\":\"draw a cat 1\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"},{\"type\":\"image_generation_call\",\"result\":\"aW1hZ2UtMg==\",\"revised_prompt\":\"draw a cat 2\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"},{\"type\":\"image_generation_call\",\"result\":\"aW1hZ2UtMw==\",\"revised_prompt\":\"draw a cat 3\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"}]}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000000,\"usage\":{\"input_tokens\":11,\"output_tokens\":22,\"input_tokens_details\":{\"cached_tokens\":3},\"output_tokens_details\":{\"image_tokens\":7}},\"tool_usage\":{\"image_gen\":{\"input_tokens\":46,\"output_tokens\":2459,\"output_tokens_details\":{\"image_tokens\":2459},\"images\":3}},\"output\":[{\"type\":\"image_generation_call\",\"result\":\"aW1hZ2UtMQ==\",\"revised_prompt\":\"draw a cat 1\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"},{\"type\":\"image_generation_call\",\"result\":\"aW1hZ2UtMg==\",\"revised_prompt\":\"draw a cat 2\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"},{\"type\":\"image_generation_call\",\"result\":\"aW1hZ2UtMw==\",\"revised_prompt\":\"draw a cat 3\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"}]}}\n\n" +
"data: [DONE]\n\n",
)),
},
@@ -655,9 +655,9 @@ func TestOpenAIGatewayServiceForwardImages_OAuthPassesNAndReturnsAllImages(t *te
require.Equal(t, "gpt-image-2", result.Model)
require.Equal(t, "gpt-image-2", result.UpstreamModel)
require.Equal(t, 3, result.ImageCount)
require.Equal(t, 11, result.Usage.InputTokens)
require.Equal(t, 22, result.Usage.OutputTokens)
require.Equal(t, 7, result.Usage.ImageOutputTokens)
require.Equal(t, 46, result.Usage.InputTokens)
require.Equal(t, 2459, result.Usage.OutputTokens)
require.Equal(t, 2459, result.Usage.ImageOutputTokens)
require.NotNil(t, upstream.lastReq)
require.Equal(t, chatgptCodexURL, upstream.lastReq.URL.String())
@@ -688,6 +688,81 @@ func TestOpenAIGatewayServiceForwardImages_OAuthPassesNAndReturnsAllImages(t *te
require.Equal(t, "draw a cat 3", gjson.Get(rec.Body.String(), "data.2.revised_prompt").String())
}
func TestParseOpenAIImagesSSEUsageBytes_ToolUsagePrecedenceAndFallback(t *testing.T) {
svc := &OpenAIGatewayService{}
fallback := OpenAIUsage{InputTokens: 3, OutputTokens: 4, ImageOutputTokens: 2}
tests := []struct {
name string
toolUsage string
want OpenAIUsage
}{
{
name: "valid tool usage takes atomic precedence",
toolUsage: `{"input_tokens":4.6e1,"output_tokens":2459e0,"output_tokens_details":{"image_tokens":24590e-1}}`,
want: OpenAIUsage{InputTokens: 46, OutputTokens: 2459, ImageOutputTokens: 2459},
},
{name: "absent", want: fallback},
{name: "malformed field", toolUsage: `{"input_tokens":"46","output_tokens":2459,"output_tokens_details":{"image_tokens":2459}}`, want: fallback},
{name: "fractional field", toolUsage: `{"input_tokens":46,"output_tokens":2459.5,"output_tokens_details":{"image_tokens":2459}}`, want: fallback},
{name: "negative field", toolUsage: `{"input_tokens":46,"output_tokens":2459,"output_tokens_details":{"image_tokens":-1}}`, want: fallback},
{name: "overflow field", toolUsage: `{"input_tokens":46,"output_tokens":9223372036854775808,"output_tokens_details":{"image_tokens":2459}}`, want: fallback},
{name: "incomplete object", toolUsage: `{"input_tokens":46,"output_tokens":2459}`, want: fallback},
{name: "hostile huge exponent", toolUsage: `{"input_tokens":1e1000000000,"output_tokens":2459,"output_tokens_details":{"image_tokens":2459}}`, want: fallback},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
toolUsageField := ""
if tt.toolUsage != "" {
toolUsageField = `,"tool_usage":{"image_gen":` + tt.toolUsage + `}`
}
payload := []byte(`{"type":"response.completed","response":{"usage":{"input_tokens":3,"output_tokens":4,"output_tokens_details":{"image_tokens":2}}` + toolUsageField + `}}`)
var got OpenAIUsage
svc.parseOpenAIImagesSSEUsageBytes(payload, &got)
require.Equal(t, tt.want, got)
})
}
}
func TestParseOpenAIImagesSSEUsageBytes_MalformedCompletedDoesNotOverrideUsage(t *testing.T) {
svc := &OpenAIGatewayService{}
var usage OpenAIUsage
svc.parseOpenAIImagesSSEUsageBytes([]byte(`{"type":"response.output_item.done","item":{"type":"image_generation_call","result":"aW1hZ2U="}}`), &usage)
svc.parseOpenAIImagesSSEUsageBytes([]byte(`{"type":"response.completed","response":{"usage":{"input_tokens":3,"output_tokens":4,"output_tokens_details":{"image_tokens":2}}}}`), &usage)
svc.parseOpenAIImagesSSEUsageBytes([]byte(`{"type":"response.completed","response":{"tool_usage":{"image_gen":{"input_tokens":46,"output_tokens":2459,"output_tokens_details":{"image_tokens":2459}}}}} trailing`), &usage)
require.Equal(t, OpenAIUsage{InputTokens: 3, OutputTokens: 4, ImageOutputTokens: 2}, usage)
}
func TestBoundedJSONNonNegativeInt(t *testing.T) {
tests := []struct {
name string
raw string
want int
ok bool
}{
{name: "scale reduction before accumulation", raw: `10000000000000000000e-19`, want: 1, ok: true},
{name: "decimal scale reduction", raw: `10000000000000000000.0e-19`, want: 1, ok: true},
{name: "fractional after scale reduction", raw: `10000000000000000001e-19`, ok: false},
{name: "overflow after scale reduction", raw: `92233720368547758080e-1`, ok: false},
{name: "zero with negative exponent", raw: `0e-100`, want: 0, ok: true},
{name: "zero beyond exponent bound", raw: `0e101`, want: 0, ok: true},
{name: "zero padded decimal beyond exponent bound", raw: `0.000000e+000000000000000000000000000000000000000000000000101`, want: 0, ok: true},
{name: "zero padded exponent", raw: `1e0000`, want: 1, ok: true},
{name: "negative zero syntax", raw: `-0e101`, ok: false},
{name: "hostile exponent", raw: `1e-1000`, ok: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, ok := boundedJSONNonNegativeInt(gjson.Parse(tt.raw))
require.Equal(t, tt.ok, ok)
require.Equal(t, tt.want, got)
})
}
}
func TestOpenAIGatewayServiceForwardImages_OAuthUpstreamHTTPErrorSurfacesRealError(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"model":"gpt-image-2","prompt":"draw a cat","response_format":"b64_json"}`)
@@ -1244,7 +1319,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthStreamingTransformsEvents(t *tes
Body: io.NopCloser(strings.NewReader(
"data: {\"type\":\"response.created\",\"response\":{\"created_at\":1710000001,\"tools\":[{\"type\":\"image_generation\",\"model\":\"gpt-image-2\",\"background\":\"auto\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"}]}}\n\n" +
"data: {\"type\":\"response.image_generation_call.partial_image\",\"partial_image_b64\":\"cGFydGlhbA==\",\"partial_image_index\":0,\"output_format\":\"png\",\"background\":\"auto\"}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000001,\"usage\":{\"input_tokens\":5,\"output_tokens\":9,\"output_tokens_details\":{\"image_tokens\":4}},\"tool_usage\":{\"image_gen\":{\"images\":1}},\"tools\":[{\"type\":\"image_generation\",\"model\":\"gpt-image-2\",\"background\":\"auto\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"}],\"output\":[{\"type\":\"image_generation_call\",\"result\":\"ZmluYWw=\",\"output_format\":\"png\"}]}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"created_at\":1710000001,\"usage\":{\"input_tokens\":5,\"output_tokens\":9,\"output_tokens_details\":{\"image_tokens\":4}},\"tool_usage\":{\"image_gen\":{\"input_tokens\":46,\"output_tokens\":2459,\"output_tokens_details\":{\"image_tokens\":2459},\"images\":1}},\"tools\":[{\"type\":\"image_generation\",\"model\":\"gpt-image-2\",\"background\":\"auto\",\"output_format\":\"png\",\"quality\":\"high\",\"size\":\"1024x1024\"}],\"output\":[{\"type\":\"image_generation_call\",\"result\":\"ZmluYWw=\",\"output_format\":\"png\"}]}}\n\n" +
"data: [DONE]\n\n",
)),
},
@@ -1266,6 +1341,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthStreamingTransformsEvents(t *tes
require.NotNil(t, result)
require.True(t, result.Stream)
require.Equal(t, 1, result.ImageCount)
require.Equal(t, OpenAIUsage{InputTokens: 46, OutputTokens: 2459, ImageOutputTokens: 2459}, result.Usage)
events := parseOpenAIImageTestSSEEvents(rec.Body.String())
partial, ok := findOpenAIImageTestSSEEvent(events, "image_generation.partial_image")
require.True(t, ok)
@@ -1290,7 +1366,7 @@ func TestOpenAIGatewayServiceForwardImages_OAuthStreamingTransformsEvents(t *tes
require.Equal(t, "high", gjson.Get(completed.Data, "quality").String())
require.Equal(t, "1024x1024", gjson.Get(completed.Data, "size").String())
require.Equal(t, "auto", gjson.Get(completed.Data, "background").String())
require.JSONEq(t, `{"images":1}`, gjson.Get(completed.Data, "usage").Raw)
require.JSONEq(t, `{"input_tokens":46,"output_tokens":2459,"output_tokens_details":{"image_tokens":2459},"images":1}`, gjson.Get(completed.Data, "usage").Raw)
require.False(t, gjson.Get(completed.Data, "revised_prompt").Exists())
}