Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 6 additions & 4 deletions relay/channel/openai/chat_via_responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
}

Expand Down
99 changes: 99 additions & 0 deletions relay/channel/openai/chat_via_responses_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,13 +11,17 @@ 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"
)

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)
Expand All @@ -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))
Comment thread
coderabbitai[bot] marked this conversation as resolved.
assert.Contains(t, recorder.Body.String(), `"usage":{"prompt_tokens":11`)
}

func TestOaiResponsesToChatStreamHandlerConvertsSSEOrderAndUsage(t *testing.T) {
oldMode := gin.Mode()
gin.SetMode(gin.TestMode)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@ type ResponsesToChatStreamState struct {
pendingArgsByOutputIndex map[int]string
pendingArgsByItemID map[string]string
usageText strings.Builder
usageTextSink func(string)
}

type responsesStreamTool struct {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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{
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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{
Expand Down
10 changes: 6 additions & 4 deletions relaykit/relayconvert/response_registry.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
Expand Down
75 changes: 75 additions & 0 deletions relaykit/relayconvert/response_registry_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package relayconvert

import (
"strings"
"testing"

"github.com/QuantumNous/new-api/relaykit/dto"
Expand Down Expand Up @@ -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{
Expand Down
Loading