diff --git a/relay/channel/openai/chat_via_responses.go b/relay/channel/openai/chat_via_responses.go index 25caeb5854bd..b84769b8a98b 100644 --- a/relay/channel/openai/chat_via_responses.go +++ b/relay/channel/openai/chat_via_responses.go @@ -194,10 +194,12 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo responseId := helper.GetResponseID(c) createAt := time.Now().Unix() + usageEst := service.NewStreamingEstimateByModel(info.UpstreamModelName) state, err := relayconvert.NewResponseStreamState(types.RelayFormatOpenAIResponses, info.RelayFormat, relayconvert.ResponseStreamOptions{ - ID: responseId, - Model: info.UpstreamModelName, - Created: createAt, + ID: responseId, + Model: info.UpstreamModelName, + Created: createAt, + UsageTextSink: usageEst.WriteString, }) if err != nil { return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError) @@ -313,7 +315,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo usage := state.Usage() if usage == nil || usage.TotalTokens == 0 { - usage = service.ResponseText2Usage(c, state.UsageText(), info.UpstreamModelName, info.GetEstimatePromptTokens()) + usage = service.StreamingEstimate2Usage(c, usageEst, info.GetEstimatePromptTokens()) state.SetUsage(usage) } diff --git a/relay/channel/openai/chat_via_responses_test.go b/relay/channel/openai/chat_via_responses_test.go index df83b1d616b9..61330edb6bc5 100644 --- a/relay/channel/openai/chat_via_responses_test.go +++ b/relay/channel/openai/chat_via_responses_test.go @@ -11,6 +11,7 @@ import ( "github.com/QuantumNous/new-api/constant" relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/relaykit/types" + "github.com/QuantumNous/new-api/service" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -18,6 +19,9 @@ import ( func newResponsesChatTestContext(t *testing.T, body string, isStream bool) (*gin.Context, *httptest.ResponseRecorder, *http.Response, *relaycommon.RelayInfo) { t.Helper() + oldTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 30 + t.Cleanup(func() { constant.StreamingTimeout = oldTimeout }) recorder := httptest.NewRecorder() c, _ := gin.CreateTestContext(recorder) @@ -39,6 +43,101 @@ func newResponsesChatTestContext(t *testing.T, body string, isStream bool) (*gin return c, recorder, resp, info } +func responsesChatSSE(events ...string) string { + lines := make([]string, 0, len(events)+2) + for _, event := range events { + lines = append(lines, "data: "+event) + } + lines = append(lines, "data: [DONE]", "") + return strings.Join(lines, "\n") +} + +func TestOaiResponsesToChatStreamFallbackUsageMatchesResponseText2Usage(t *testing.T) { + body := responsesChatSSE( + `{"type":"response.created","response":{"id":"resp_test","created_at":1710000000,"model":"gpt-test"}}`, + `{"type":"response.reasoning_summary_text.delta","delta":"Reasoning summary 123"}`, + `{"type":"response.reasoning_summary_text.done"}`, + `{"type":"response.reasoning_summary_text.delta","delta":"second paragraph"}`, + `{"type":"response.output_text.delta","delta":" Visible output 中文"}`, + `{"type":"response.completed","response":{"usage":{"input_tokens":0,"output_tokens":0,"total_tokens":0}}}`, + ) + + c, recorder, resp, info := newResponsesChatTestContext(t, body, true) + info.SetEstimatePromptTokens(37) + usage, err := OaiResponsesToChatStreamHandler(c, info, resp) + require.Nil(t, err) + require.NotNil(t, usage) + + expectedContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + expected := service.ResponseText2Usage(expectedContext, "Reasoning summary 123\n\nsecond paragraph Visible output 中文", info.UpstreamModelName, info.GetEstimatePromptTokens()) + assert.Equal(t, expected, usage) + assert.True(t, common.GetContextKeyBool(c, constant.ContextKeyLocalCountTokens)) + assert.Contains(t, recorder.Body.String(), `"usage":{"prompt_tokens":37`) +} + +func TestOaiResponsesToChatStreamFallbackUsageCountsToolCalls(t *testing.T) { + body := responsesChatSSE( + `{"type":"response.created","response":{"id":"resp_test","created_at":1710000000,"model":"gpt-test"}}`, + `{"type":"response.output_item.added","output_index":0,"item":{"id":"fc_1","type":"function_call","call_id":"call_1","name":"lookup"}}`, + `{"type":"response.function_call_arguments.delta","output_index":0,"item_id":"fc_1","delta":"{\"city\":\"Beijing\"}"}`, + `{"type":"response.completed","response":{"usage":{"input_tokens":0,"output_tokens":0,"total_tokens":0}}}`, + ) + + c, recorder, resp, info := newResponsesChatTestContext(t, body, true) + info.SetEstimatePromptTokens(37) + usage, err := OaiResponsesToChatStreamHandler(c, info, resp) + require.Nil(t, err) + require.NotNil(t, usage) + + expectedContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + expected := service.ResponseText2Usage(expectedContext, `lookup{"city":"Beijing"}`, info.UpstreamModelName, info.GetEstimatePromptTokens()) + assert.Equal(t, expected, usage) + assert.True(t, common.GetContextKeyBool(c, constant.ContextKeyLocalCountTokens)) + assert.Contains(t, recorder.Body.String(), `"finish_reason":"tool_calls"`) +} + +func TestOaiResponsesToChatStreamUsesUpstreamUsageWithoutLocalCountFlag(t *testing.T) { + body := responsesChatSSE( + `{"type":"response.created","response":{"id":"resp_test","created_at":1710000000,"model":"gpt-test"}}`, + `{"type":"response.output_text.delta","delta":"Visible output that must not be locally counted"}`, + `{"type":"response.completed","response":{"usage":{"input_tokens":11,"output_tokens":13,"total_tokens":24,"input_tokens_details":{"cached_tokens":7},"completion_tokens_details":{"reasoning_tokens":5}}}}`, + ) + + c, recorder, resp, info := newResponsesChatTestContext(t, body, true) + usage, err := OaiResponsesToChatStreamHandler(c, info, resp) + require.Nil(t, err) + require.NotNil(t, usage) + assert.Equal(t, 11, usage.PromptTokens) + assert.Equal(t, 13, usage.CompletionTokens) + assert.Equal(t, 24, usage.TotalTokens) + assert.Equal(t, 11, usage.InputTokens) + assert.Equal(t, 13, usage.OutputTokens) + assert.Equal(t, 7, usage.PromptTokensDetails.CachedTokens) + assert.Equal(t, 5, usage.CompletionTokenDetails.ReasoningTokens) + assert.False(t, common.GetContextKeyBool(c, constant.ContextKeyLocalCountTokens)) + assert.Contains(t, recorder.Body.String(), `"usage":{"prompt_tokens":11`) +} + +func TestOaiResponsesToChatStreamCompletesUpstreamUsageTotalWithoutLocalCountFlag(t *testing.T) { + body := responsesChatSSE( + `{"type":"response.created","response":{"id":"resp_test","created_at":1710000000,"model":"gpt-test"}}`, + `{"type":"response.output_text.delta","delta":"Visible output with upstream input and output tokens"}`, + `{"type":"response.completed","response":{"usage":{"input_tokens":11,"output_tokens":13,"total_tokens":0}}}`, + ) + + c, recorder, resp, info := newResponsesChatTestContext(t, body, true) + usage, err := OaiResponsesToChatStreamHandler(c, info, resp) + require.Nil(t, err) + require.NotNil(t, usage) + assert.Equal(t, 11, usage.PromptTokens) + assert.Equal(t, 13, usage.CompletionTokens) + assert.Equal(t, 24, usage.TotalTokens) + assert.Equal(t, 11, usage.InputTokens) + assert.Equal(t, 13, usage.OutputTokens) + assert.False(t, common.GetContextKeyBool(c, constant.ContextKeyLocalCountTokens)) + assert.Contains(t, recorder.Body.String(), `"usage":{"prompt_tokens":11`) +} + func TestOaiResponsesToChatStreamHandlerConvertsSSEOrderAndUsage(t *testing.T) { oldMode := gin.Mode() gin.SetMode(gin.TestMode) diff --git a/relaykit/relayconvert/internal/oai_responses/to_oai_chat_stream_resp.go b/relaykit/relayconvert/internal/oai_responses/to_oai_chat_stream_resp.go index 6026e3899eeb..39e955fdabde 100644 --- a/relaykit/relayconvert/internal/oai_responses/to_oai_chat_stream_resp.go +++ b/relaykit/relayconvert/internal/oai_responses/to_oai_chat_stream_resp.go @@ -32,6 +32,7 @@ type ResponsesToChatStreamState struct { pendingArgsByOutputIndex map[int]string pendingArgsByItemID map[string]string usageText strings.Builder + usageTextSink func(string) } type responsesStreamTool struct { @@ -68,6 +69,24 @@ func (s *ResponsesToChatStreamState) UsageText() string { return s.usageText.String() } +func (s *ResponsesToChatStreamState) SetUsageTextSink(sink func(string)) { + if s == nil { + return + } + s.usageTextSink = sink +} + +func (s *ResponsesToChatStreamState) recordUsageText(text string) { + if s == nil || text == "" { + return + } + if s.usageTextSink != nil { + s.usageTextSink(text) + return + } + s.usageText.WriteString(text) +} + func ResponsesStreamEventToChatChunks(event *dto.ResponsesStreamResponse, state *ResponsesToChatStreamState) ([]dto.ChatCompletionsStreamResponse, error) { if event == nil || state == nil { return nil, nil @@ -151,7 +170,7 @@ func (s *ResponsesToChatStreamState) textDelta(delta string) []dto.ChatCompletio if delta == "" { return nil } - s.usageText.WriteString(delta) + s.recordUsageText(delta) s.hasSentText = true chunks := s.ensureStart() chunks = append(chunks, s.makeChunk(dto.ChatCompletionsStreamResponseChoiceDelta{ @@ -207,7 +226,7 @@ func (s *ResponsesToChatStreamState) reasoningDelta(delta string) []dto.ChatComp s.needsReasoningSummaryBreak = false } } - s.usageText.WriteString(delta) + s.recordUsageText(delta) chunks := s.ensureStart() chunks = append(chunks, s.makeChunk(dto.ChatCompletionsStreamResponseChoiceDelta{ ReasoningContent: &delta, @@ -418,10 +437,10 @@ func (s *ResponsesToChatStreamState) toolDelta(tool *responsesStreamTool, explic } if argsDelta != "" { tool.ArgsSentAt += len(argsDelta) - s.usageText.WriteString(argsDelta) + s.recordUsageText(argsDelta) } if responseTool.Function.Name != "" { - s.usageText.WriteString(responseTool.Function.Name) + s.recordUsageText(responseTool.Function.Name) } chunks = append(chunks, s.makeChunk(dto.ChatCompletionsStreamResponseChoiceDelta{ diff --git a/relaykit/relayconvert/response_registry.go b/relaykit/relayconvert/response_registry.go index a2369a61eda1..45945012fd1d 100644 --- a/relaykit/relayconvert/response_registry.go +++ b/relaykit/relayconvert/response_registry.go @@ -70,10 +70,11 @@ type responseConverterRoute struct { } type ResponseStreamOptions struct { - ID string - Model string - Created int64 - IncludeUsage bool + ID string + Model string + Created int64 + IncludeUsage bool + UsageTextSink func(string) } type ResponseStreamState struct { @@ -855,6 +856,7 @@ func finalizeOAIChatStreamResponseToOAIResponses(_ context.Context, _ convmeta.M func newOAIResponsesToOAIChatStreamState(options ResponseStreamOptions) any { state := NewResponsesToChatStreamState(strings.TrimSpace(options.Model), options.IncludeUsage) + state.SetUsageTextSink(options.UsageTextSink) state.ID = strings.TrimSpace(options.ID) if options.Created != 0 { state.Created = options.Created diff --git a/relaykit/relayconvert/response_registry_test.go b/relaykit/relayconvert/response_registry_test.go index 3e62d4c2ce4c..c767415a74e4 100644 --- a/relaykit/relayconvert/response_registry_test.go +++ b/relaykit/relayconvert/response_registry_test.go @@ -1,6 +1,7 @@ package relayconvert import ( + "strings" "testing" "github.com/QuantumNous/new-api/relaykit/dto" @@ -489,6 +490,80 @@ func TestConvertStreamResponseStatefulDirectConverters(t *testing.T) { require.IsType(t, dto.ChatCompletionsStreamResponse{}, responsesResults[len(responsesResults)-1].Value) } +func TestResponseStreamUsageTextSinkMatchesRetainedUsageText(t *testing.T) { + events := []*dto.ResponsesStreamResponse{ + {Type: "response.reasoning_summary_text.delta", Delta: "Reasoning summary"}, + {Type: "response.reasoning_summary_text.done"}, + {Type: "response.reasoning_summary_text.delta", Delta: "second paragraph"}, + {Type: "response.output_text.delta", Delta: " visible output"}, + { + Type: "response.output_item.added", + OutputIndex: respPtr(0), + Item: &dto.ResponsesOutput{ + Type: "function_call", + ID: "fc_1", + CallId: "call_1", + Name: "lookup", + Arguments: []byte(`{"city":"Bei`), + }, + }, + {Type: "response.function_call_arguments.delta", OutputIndex: respPtr(0), ItemID: "fc_1", Delta: `jing"}`}, + } + wantChunks := []string{ + "Reasoning summary", + "\n\nsecond paragraph", + " visible output", + `{"city":"Bei`, + "lookup", + `jing"}`, + } + + var gotChunks []string + sinkState, err := NewResponseStreamState(types.RelayFormatOpenAIResponses, types.RelayFormatOpenAI, ResponseStreamOptions{ + Model: "gpt-test", + UsageTextSink: func(text string) { + gotChunks = append(gotChunks, text) + }, + }) + require.NoError(t, err) + for _, event := range events { + _, err = ConvertStreamResponseChunk(nil, nil, sinkState, event) + require.NoError(t, err) + } + assert.Equal(t, wantChunks, gotChunks) + assert.Empty(t, sinkState.UsageText()) + + retainedState, err := NewResponseStreamState(types.RelayFormatOpenAIResponses, types.RelayFormatOpenAI, ResponseStreamOptions{Model: "gpt-test"}) + require.NoError(t, err) + for _, event := range events { + _, err = ConvertStreamResponseChunk(nil, nil, retainedState, event) + require.NoError(t, err) + } + assert.Equal(t, strings.Join(wantChunks, ""), retainedState.UsageText()) +} + +func TestResponseStreamUsageTextSinkRunsOnceInMultiHopRoute(t *testing.T) { + var got []string + state, err := NewResponseStreamState(types.RelayFormatOpenAIResponses, types.RelayFormatClaude, ResponseStreamOptions{ + Model: "gpt-test", + UsageTextSink: func(text string) { + got = append(got, text) + }, + }) + require.NoError(t, err) + info := &convmeta.Values{ + ClaudeConvertInfo: &convmeta.ClaudeConvertInfo{LastMessagesType: convmeta.LastMessageTypeNone}, + } + + _, err = ConvertStreamResponseChunk(nil, info, state, &dto.ResponsesStreamResponse{ + Type: "response.output_text.delta", + Delta: "hello", + }) + require.NoError(t, err) + assert.Equal(t, []string{"hello"}, got) + assert.Empty(t, state.UsageText()) +} + func TestConvertStreamResponseStatefulMultiHopResponsesToClaude(t *testing.T) { info := &convmeta.Values{ ClaudeConvertInfo: &convmeta.ClaudeConvertInfo{ diff --git a/service/stream_estimate.go b/service/stream_estimate.go new file mode 100644 index 000000000000..a07618bb7d9d --- /dev/null +++ b/service/stream_estimate.go @@ -0,0 +1,167 @@ +package service + +import ( + "math" + "strings" + "unicode" + "unicode/utf8" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/relaykit/dto" + "github.com/gin-gonic/gin" +) + +type streamingEstimateWordType int + +const ( + streamingEstimateNone streamingEstimateWordType = iota + streamingEstimateLatin + streamingEstimateNumber +) + +type StreamingEstimateByModel struct { + m multipliers + count float64 + currentWordType streamingEstimateWordType + pending []byte + hasInput bool +} + +func NewStreamingEstimateByModel(model string) *StreamingEstimateByModel { + model = strings.ToLower(model) + provider := OpenAI + if strings.Contains(model, "gemini") { + provider = Gemini + } else if strings.Contains(model, "claude") { + provider = Claude + } + return &StreamingEstimateByModel{m: getMultipliers(provider)} +} + +func (e *StreamingEstimateByModel) WriteString(text string) { + if e == nil || text == "" { + return + } + e.hasInput = true + if len(e.pending) > 0 { + for len(text) > 0 { + e.pending = append(e.pending, text[0]) + text = text[1:] + if utf8.FullRune(e.pending) { + break + } + } + if !utf8.FullRune(e.pending) { + return + } + r, size := utf8.DecodeRune(e.pending) + e.writeRune(r) + if size < len(e.pending) { + text = string(e.pending[size:]) + text + } + e.pending = e.pending[:0] + } + prefix, pending := splitTrailingIncompleteUTF8(text) + for _, r := range prefix { + e.writeRune(r) + } + if pending != "" { + e.pending = append(e.pending, pending...) + } +} + +func (e *StreamingEstimateByModel) Tokens() int { + if e == nil { + return 0 + } + snapshot := *e + if len(snapshot.pending) > 0 { + for _, r := range string(snapshot.pending) { + snapshot.writeRune(r) + } + } + if !snapshot.hasInput { + return 0 + } + return int(math.Ceil(snapshot.count)) + snapshot.m.BasePad +} + +func StreamingEstimate2Usage(c *gin.Context, e *StreamingEstimateByModel, promptTokens int) *dto.Usage { + common.SetContextKey(c, constant.ContextKeyLocalCountTokens, true) + usage := &dto.Usage{ + PromptTokens: promptTokens, + CompletionTokens: e.Tokens(), + } + usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens + return usage +} + +func (e *StreamingEstimateByModel) writeRune(r rune) { + if unicode.IsSpace(r) { + e.currentWordType = streamingEstimateNone + if r == '\n' || r == '\t' { + e.count += e.m.Newline + } else { + e.count += e.m.Space + } + return + } + + if isCJK(r) { + e.currentWordType = streamingEstimateNone + e.count += e.m.CJK + return + } + + if isEmoji(r) { + e.currentWordType = streamingEstimateNone + e.count += e.m.Emoji + return + } + + if isLatinOrNumber(r) { + newType := streamingEstimateLatin + if unicode.IsNumber(r) { + newType = streamingEstimateNumber + } + if e.currentWordType == streamingEstimateNone || e.currentWordType != newType { + if newType == streamingEstimateNumber { + e.count += e.m.Number + } else { + e.count += e.m.Word + } + e.currentWordType = newType + } + return + } + + e.currentWordType = streamingEstimateNone + if isMathSymbol(r) { + e.count += e.m.MathSymbol + } else if r == '@' { + e.count += e.m.AtSign + } else if isURLDelim(r) { + e.count += e.m.URLDelim + } else { + e.count += e.m.Symbol + } +} + +func splitTrailingIncompleteUTF8(text string) (string, string) { + if text == "" { + return "", "" + } + start := len(text) - 1 + for start > 0 && isUTF8Continuation(text[start]) { + start-- + } + if start < len(text) && !utf8.FullRuneInString(text[start:]) { + return text[:start], text[start:] + } + return text, "" +} + +func isUTF8Continuation(b byte) bool { + return b&0xc0 == 0x80 +} diff --git a/service/stream_estimate_test.go b/service/stream_estimate_test.go new file mode 100644 index 000000000000..8edea3770b3f --- /dev/null +++ b/service/stream_estimate_test.go @@ -0,0 +1,148 @@ +package service + +import ( + "net/http/httptest" + "strings" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestStreamingEstimateByModelMatchesEstimateTokenByModel(t *testing.T) { + cases := []struct { + name string + model string + chunks []string + }{ + { + name: "openai ascii split word", + model: "gpt-5.5", + chunks: []string{ + "Deterministic account", + "ing validation with 123", + "456 and https://example.com/a?b=c", + }, + }, + { + name: "claude cjk and symbols", + model: "claude-opus-4-8", + chunks: []string{ + "中文混合 English ", + "∑∫√ and emoji ✅", + "\nnew line", + }, + }, + { + name: "gemini long mixed", + model: "gemini-2.5-pro", + chunks: []string{ + strings.Repeat("alpha123 中文 ", 100), + strings.Repeat(" /path?x=y&z=1\n", 100), + }, + }, + { + name: "empty text remains zero", + model: "gpt-5.5", + chunks: nil, + }, + { + name: "whitespace only", + model: "gpt-5.5", + chunks: []string{ + " ", + "\n\t", + " ", + }, + }, + { + name: "symbols and mixed word boundaries", + model: "claude-opus-4-8", + chunks: []string{ + "abc", + "123", + "xyz@example.com", + " ∑∫√∞ /:?&=;#%", + }, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + joined := strings.Join(tc.chunks, "") + estimator := NewStreamingEstimateByModel(tc.model) + for _, chunk := range tc.chunks { + estimator.WriteString(chunk) + } + got := estimator.Tokens() + want := EstimateTokenByModel(tc.model, joined) + require.Equal(t, want, got) + }) + } +} + +func TestStreamingEstimateByModelMatchesSplitUTF8(t *testing.T) { + text := "emoji ✅ 中文 ∑ math and URL https://example.com/a?b=c" + data := []byte(text) + estimator := NewStreamingEstimateByModel("claude-opus-4-8") + for i := 0; i < len(data); i++ { + estimator.WriteString(string(data[i : i+1])) + } + got := estimator.Tokens() + want := EstimateTokenByModel("claude-opus-4-8", text) + require.Equal(t, want, got) +} + +func TestStreamingEstimateByModelMatchesTrailingInvalidUTF8(t *testing.T) { + chunks := []string{"valid 中文 ", string([]byte{0xe2, 0x82})} + text := strings.Join(chunks, "") + estimator := NewStreamingEstimateByModel("gpt-5.5") + for _, chunk := range chunks { + estimator.WriteString(chunk) + } + got := estimator.Tokens() + want := EstimateTokenByModel("gpt-5.5", text) + require.Equal(t, want, got) +} + +func TestStreamingEstimateByModelMatchesPendingPlusMoreText(t *testing.T) { + data := []byte("✅abc123中文") + chunks := []string{ + string(data[:1]), + string(data[1:5]), + string(data[5:]), + } + text := strings.Join(chunks, "") + estimator := NewStreamingEstimateByModel("gemini-2.5-pro") + for _, chunk := range chunks { + estimator.WriteString(chunk) + } + got := estimator.Tokens() + want := EstimateTokenByModel("gemini-2.5-pro", text) + require.Equal(t, want, got) +} + +func TestStreamingEstimateToUsageMatchesResponseText2Usage(t *testing.T) { + text := "Reasoning summary 123\n\nVisible output 中文 with https://example.com/a?b=c" + model := "gpt-5.5" + promptTokens := 37 + + streaming := NewStreamingEstimateByModel(model) + for _, chunk := range []string{"Reasoning summary 123", "\n\nVisible output 中文", " with https://example.com/a?b=c"} { + streaming.WriteString(chunk) + } + + gotContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + got := StreamingEstimate2Usage(gotContext, streaming, promptTokens) + + wantContext, _ := gin.CreateTestContext(httptest.NewRecorder()) + want := ResponseText2Usage(wantContext, text, model, promptTokens) + + assert.Equal(t, want.PromptTokens, got.PromptTokens) + assert.Equal(t, want.CompletionTokens, got.CompletionTokens) + assert.Equal(t, want.TotalTokens, got.TotalTokens) + assert.True(t, common.GetContextKeyBool(gotContext, constant.ContextKeyLocalCountTokens)) +} diff --git a/service/usage_count_bench_test.go b/service/usage_count_bench_test.go new file mode 100644 index 000000000000..b178e0e5bd70 --- /dev/null +++ b/service/usage_count_bench_test.go @@ -0,0 +1,146 @@ +package service + +import ( + "net/http/httptest" + "strings" + "testing" + + "github.com/gin-gonic/gin" +) + +var usageCountBenchSink int + +func init() { + InitTokenEncoders() + gin.SetMode(gin.TestMode) +} + +func benchmarkText(repeat int) string { + parts := []string{ + "Deterministic accounting validation compares prompt tokens, completion tokens, quota, stream status, and request identifiers.", + "中文内容用于覆盖 CJK 计数路径,避免只测英文 ASCII 文本。", + "Numbers 1234567890, URLs https://example.com/a?b=c&d=e, symbols ∑∫√∞, and emoji ✅ are included.", + } + return strings.Repeat(strings.Join(parts, "\n"), repeat) +} + +func BenchmarkEstimateTokenByModel(b *testing.B) { + cases := []struct { + name string + model string + repeat int + }{ + {name: "openai_small", model: "gpt-5.5", repeat: 1}, + {name: "openai_large", model: "gpt-5.5", repeat: 400}, + {name: "claude_small", model: "claude-sonnet-4", repeat: 1}, + {name: "claude_large", model: "claude-sonnet-4", repeat: 400}, + } + + for _, tc := range cases { + text := benchmarkText(tc.repeat) + b.Run(tc.name, func(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + usageCountBenchSink = EstimateTokenByModel(tc.model, text) + } + }) + } +} + +func BenchmarkCountTextToken(b *testing.B) { + cases := []struct { + name string + model string + repeat int + }{ + {name: "openai_small", model: "gpt-5.5", repeat: 1}, + {name: "openai_large", model: "gpt-5.5", repeat: 400}, + {name: "claude_small", model: "claude-sonnet-4", repeat: 1}, + {name: "claude_large", model: "claude-sonnet-4", repeat: 400}, + } + + for _, tc := range cases { + text := benchmarkText(tc.repeat) + b.Run(tc.name, func(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + usageCountBenchSink = CountTextToken(text, tc.model) + } + }) + } +} + +func BenchmarkResponseText2Usage(b *testing.B) { + cases := []struct { + name string + model string + repeat int + }{ + {name: "openai_small", model: "gpt-5.5", repeat: 1}, + {name: "openai_large", model: "gpt-5.5", repeat: 400}, + {name: "claude_small", model: "claude-sonnet-4", repeat: 1}, + {name: "claude_large", model: "claude-sonnet-4", repeat: 400}, + } + + for _, tc := range cases { + text := benchmarkText(tc.repeat) + b.Run(tc.name, func(b *testing.B) { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + usage := ResponseText2Usage(c, text, tc.model, 123) + usageCountBenchSink = usage.TotalTokens + } + }) + } +} + +func BenchmarkChatViaResponsesFallbackUsage(b *testing.B) { + chunks := []string{ + "Reasoning summary 123", + "\n\nsecond paragraph", + " Visible output 中文", + `lookup{"city":"Beijing"}`, + " with https://example.com/a?b=c&d=e and symbols ∑∫√∞ ✅", + } + largeChunks := make([]string, 0, len(chunks)*400) + for i := 0; i < 400; i++ { + largeChunks = append(largeChunks, chunks...) + } + + cases := []struct { + name string + model string + chunks []string + }{ + {name: "old_builder_openai_large", model: "gpt-5.5", chunks: largeChunks}, + {name: "new_streaming_openai_large", model: "gpt-5.5", chunks: largeChunks}, + {name: "old_builder_claude_large", model: "claude-sonnet-4", chunks: largeChunks}, + {name: "new_streaming_claude_large", model: "claude-sonnet-4", chunks: largeChunks}, + } + + for _, tc := range cases { + b.Run(tc.name, func(b *testing.B) { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + b.ReportAllocs() + for i := 0; i < b.N; i++ { + if strings.HasPrefix(tc.name, "old_builder") { + var usageText strings.Builder + for _, chunk := range tc.chunks { + usageText.WriteString(chunk) + } + usage := ResponseText2Usage(c, usageText.String(), tc.model, 123) + usageCountBenchSink = usage.TotalTokens + continue + } + + estimator := NewStreamingEstimateByModel(tc.model) + for _, chunk := range tc.chunks { + estimator.WriteString(chunk) + } + usage := StreamingEstimate2Usage(c, estimator, 123) + usageCountBenchSink = usage.TotalTokens + } + }) + } +}