mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-21 14:19:18 +08:00
fix(openai): use image tool usage for OAuth billing
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user