diff --git a/backend/internal/service/openai_images_responses.go b/backend/internal/service/openai_images_responses.go index 151de2c4b1..f5b99ac9ca 100644 --- a/backend/internal/service/openai_images_responses.go +++ b/backend/internal/service/openai_images_responses.go @@ -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 } diff --git a/backend/internal/service/openai_images_test.go b/backend/internal/service/openai_images_test.go index 0bcd68a386..a76fa414a3 100644 --- a/backend/internal/service/openai_images_test.go +++ b/backend/internal/service/openai_images_test.go @@ -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()) }