diff --git a/relay/channel/openai/chat_via_responses_test.go b/relay/channel/openai/chat_via_responses_test.go index 2bf98b56..df83b1d6 100644 --- a/relay/channel/openai/chat_via_responses_test.go +++ b/relay/channel/openai/chat_via_responses_test.go @@ -12,6 +12,7 @@ import ( relaycommon "github.com/QuantumNous/new-api/relay/common" "github.com/QuantumNous/new-api/relaykit/types" "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -86,6 +87,60 @@ func TestOaiResponsesToChatStreamHandlerConvertsSSEOrderAndUsage(t *testing.T) { ) } +func TestOaiResponsesToChatStreamHandlerConvertsClaudeSSETerminalsAndUsage(t *testing.T) { + oldMode := gin.Mode() + gin.SetMode(gin.TestMode) + t.Cleanup(func() { gin.SetMode(oldMode) }) + + oldTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 30 + t.Cleanup(func() { constant.StreamingTimeout = oldTimeout }) + + body := strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_1","model":"gpt-test","created_at":1710000000}}`, + `data: {"type":"response.output_text.delta","delta":"hello"}`, + `data: {"type":"response.completed","response":{"status":"completed","usage":{"input_tokens":2,"output_tokens":3,"total_tokens":5}}}`, + `data: [DONE]`, + ``, + }, "\n") + + c, recorder, resp, info := newResponsesChatTestContext(t, body, true) + info.RelayFormat = types.RelayFormatClaude + + usage, err := OaiResponsesToChatStreamHandler(c, info, resp) + require.Nil(t, err) + require.NotNil(t, usage) + assert.Equal(t, 2, usage.PromptTokens) + assert.Equal(t, 3, usage.CompletionTokens) + assert.Equal(t, 5, usage.TotalTokens) + + got := recorder.Body.String() + assert.Equal(t, "text/event-stream", recorder.Header().Get("Content-Type")) + assert.Equal(t, 1, strings.Count(got, "event: message_start\n")) + assert.Equal(t, 1, strings.Count(got, "event: content_block_stop\n")) + assert.Equal(t, 1, strings.Count(got, "event: message_delta\n")) + assert.Equal(t, 1, strings.Count(got, "event: message_stop\n")) + + messageDeltaFrame := "" + for _, frame := range strings.Split(got, "\n\n") { + if strings.HasPrefix(frame, "event: message_delta\n") { + messageDeltaFrame = frame + break + } + } + require.NotEmpty(t, messageDeltaFrame) + assert.Contains(t, messageDeltaFrame, `"type":"message_delta"`) + assert.Contains(t, messageDeltaFrame, `"stop_reason":"end_turn"`) + assert.Contains(t, messageDeltaFrame, `"input_tokens":2`) + assert.Contains(t, messageDeltaFrame, `"output_tokens":3`) + requireOrderedSubstrings(t, got, + "event: message_start\n", + "event: content_block_stop\n", + "event: message_delta\n", + "event: message_stop\n", + ) +} + func TestOaiResponsesToChatBufferedStreamHandlerReturnsJSONFromSSE(t *testing.T) { oldMode := gin.Mode() gin.SetMode(gin.TestMode) diff --git a/relaykit/relayconvert/internal/gemini_chat/to_oai_chat_resp.go b/relaykit/relayconvert/internal/gemini_chat/to_oai_chat_resp.go index e256fd10..81f8078d 100644 --- a/relaykit/relayconvert/internal/gemini_chat/to_oai_chat_resp.go +++ b/relaykit/relayconvert/internal/gemini_chat/to_oai_chat_resp.go @@ -282,6 +282,107 @@ func StreamResponseGeminiChat2OpenAI(geminiResponse *dto.GeminiChatResponse) (*d return &response, isStop } +type GeminiToChatStreamState struct { + id string + created int64 + sawToolCall bool + finishEmitted bool + latestUsage *dto.Usage +} + +func NewGeminiToChatStreamState(id string, created int64) *GeminiToChatStreamState { + id = strings.TrimSpace(id) + if id == "" { + id = fmt.Sprintf("chatcmpl-%s", kitutil.GetUUID()) + } + if created == 0 { + created = kitutil.GetTimestamp() + } + return &GeminiToChatStreamState{id: id, created: created} +} + +func (s *GeminiToChatStreamState) ConvertChunk(geminiResponse *dto.GeminiChatResponse, model string, usage *dto.Usage) []*dto.ChatCompletionsStreamResponse { + if s == nil || geminiResponse == nil { + return nil + } + hasNonStopFinish := false + for _, candidate := range geminiResponse.Candidates { + if candidate.FinishReason != nil && *candidate.FinishReason != "" && *candidate.FinishReason != "STOP" { + hasNonStopFinish = true + break + } + } + response, isStop := StreamResponseGeminiChat2OpenAI(geminiResponse) + if response == nil { + return nil + } + response.Id = s.id + response.Created = s.created + response.Model = model + response.Usage = usage + + if response.IsToolCall() { + s.sawToolCall = true + if !hasNonStopFinish { + for i := range response.Choices { + if response.Choices[i].FinishReason != nil && *response.Choices[i].FinishReason == types.FinishReasonToolCalls { + response.Choices[i].FinishReason = nil + } + } + } + } + if usage != nil { + s.latestUsage = usage + } + for _, choice := range response.Choices { + if choice.FinishReason != nil && *choice.FinishReason != "" { + s.finishEmitted = true + break + } + } + + responses := []*dto.ChatCompletionsStreamResponse{response} + if isStop && !s.finishEmitted { + responses = append(responses, s.terminalChunk(model)) + } + return responses +} + +func (s *GeminiToChatStreamState) Finalize(model string) []*dto.ChatCompletionsStreamResponse { + if s == nil || s.finishEmitted { + return nil + } + return []*dto.ChatCompletionsStreamResponse{s.terminalChunk(model)} +} + +func (s *GeminiToChatStreamState) Usage() *dto.Usage { + if s == nil { + return nil + } + return s.latestUsage +} + +func (s *GeminiToChatStreamState) terminalChunk(model string) *dto.ChatCompletionsStreamResponse { + finishReason := types.FinishReasonStop + if s.sawToolCall { + finishReason = types.FinishReasonToolCalls + } + s.finishEmitted = true + return &dto.ChatCompletionsStreamResponse{ + Id: s.id, + Object: "chat.completion.chunk", + Created: s.created, + Model: model, + Choices: []dto.ChatCompletionsStreamResponseChoice{ + { + Delta: dto.ChatCompletionsStreamResponseChoiceDelta{}, + FinishReason: &finishReason, + }, + }, + Usage: s.latestUsage, + } +} + func geminiResponseToolCall(item *dto.GeminiPart) *dto.ToolCallResponse { argsBytes, err := kitutil.Marshal(item.FunctionCall.Arguments) if err != nil { diff --git a/relaykit/relayconvert/internal/oai_chat/to_claude_messages_resp.go b/relaykit/relayconvert/internal/oai_chat/to_claude_messages_resp.go index ff958b0b..554e70a9 100644 --- a/relaykit/relayconvert/internal/oai_chat/to_claude_messages_resp.go +++ b/relaykit/relayconvert/internal/oai_chat/to_claude_messages_resp.go @@ -17,6 +17,24 @@ func generateStopBlock(index int) *dto.ClaudeResponse { } } +func stopOpenBlocks(state *convmeta.ClaudeConvertInfo) []*dto.ClaudeResponse { + if state == nil { + return nil + } + switch state.LastMessagesType { + case convmeta.LastMessageTypeText, convmeta.LastMessageTypeThinking: + return []*dto.ClaudeResponse{generateStopBlock(state.Index)} + case convmeta.LastMessageTypeTools: + responses := make([]*dto.ClaudeResponse, 0, state.ToolCallMaxIndexOffset+1) + for offset := 0; offset <= state.ToolCallMaxIndexOffset; offset++ { + responses = append(responses, generateStopBlock(state.ToolCallBaseIndex+offset)) + } + return responses + default: + return nil + } +} + func buildClaudeUsageFromOpenAIUsage(oaiUsage *dto.Usage) *dto.ClaudeUsage { if oaiUsage == nil { return nil @@ -89,16 +107,8 @@ func StreamResponseOpenAI2Claude(openAIResponse *dto.ChatCompletionsStreamRespon // For text/thinking, there is at most one open block at state.Index. // For tools, OpenAI tool_calls can stream multiple parallel tool_use blocks (indexed from 0), // so we may have multiple open blocks and must stop each one explicitly. - stopOpenBlocks := func() { - switch state.LastMessagesType { - case convmeta.LastMessageTypeText, convmeta.LastMessageTypeThinking: - claudeResponses = append(claudeResponses, generateStopBlock(state.Index)) - case convmeta.LastMessageTypeTools: - base := state.ToolCallBaseIndex - for offset := 0; offset <= state.ToolCallMaxIndexOffset; offset++ { - claudeResponses = append(claudeResponses, generateStopBlock(base+offset)) - } - } + appendStopOpenBlocks := func() { + claudeResponses = append(claudeResponses, stopOpenBlocks(state)...) } // stopOpenBlocksAndAdvance closes the currently open block(s) and advances the content block index // to the next available slot for subsequent content_block_start events. @@ -109,7 +119,7 @@ func StreamResponseOpenAI2Claude(openAIResponse *dto.ChatCompletionsStreamRespon if state.LastMessagesType == convmeta.LastMessageTypeNone { return } - stopOpenBlocks() + appendStopOpenBlocks() switch state.LastMessagesType { case convmeta.LastMessageTypeTools: state.Index = state.ToolCallBaseIndex + state.ToolCallMaxIndexOffset + 1 @@ -234,23 +244,24 @@ func StreamResponseOpenAI2Claude(openAIResponse *dto.ChatCompletionsStreamRespon } } - // 如果首块就带 finish_reason,需要立即发送停止块 + // A first chunk can carry finish_reason before usage; defer terminal events until usage arrives. if len(openAIResponse.Choices) > 0 && openAIResponse.Choices[0].FinishReason != nil && *openAIResponse.Choices[0].FinishReason != "" { state.FinishReason = *openAIResponse.Choices[0].FinishReason - stopOpenBlocks() oaiUsage := openAIResponse.Usage if oaiUsage == nil { oaiUsage = state.Usage } - if oaiUsage != nil { - claudeResponses = append(claudeResponses, &dto.ClaudeResponse{ - Type: "message_delta", - Usage: buildClaudeUsageFromOpenAIUsage(oaiUsage), - Delta: &dto.ClaudeMediaMessage{ - StopReason: kitutil.GetPointer[string](stopReasonOpenAI2Claude(state.FinishReason)), - }, - }) + if oaiUsage == nil { + return claudeResponses } + appendStopOpenBlocks() + claudeResponses = append(claudeResponses, &dto.ClaudeResponse{ + Type: "message_delta", + Usage: buildClaudeUsageFromOpenAIUsage(oaiUsage), + Delta: &dto.ClaudeMediaMessage{ + StopReason: kitutil.GetPointer[string](stopReasonOpenAI2Claude(state.FinishReason)), + }, + }) claudeResponses = append(claudeResponses, &dto.ClaudeResponse{ Type: "message_stop", }) @@ -266,7 +277,7 @@ func StreamResponseOpenAI2Claude(openAIResponse *dto.ChatCompletionsStreamRespon oaiUsage = state.Usage } if oaiUsage != nil { - stopOpenBlocks() + appendStopOpenBlocks() stopReason := stopReasonOpenAI2Claude(state.FinishReason) if stopReason == "" { stopReason = "end_turn" @@ -403,7 +414,7 @@ func StreamResponseOpenAI2Claude(openAIResponse *dto.ChatCompletionsStreamRespon } if doneChunk || state.Done { - stopOpenBlocks() + appendStopOpenBlocks() oaiUsage := openAIResponse.Usage if oaiUsage == nil { oaiUsage = state.Usage @@ -428,6 +439,34 @@ func StreamResponseOpenAI2Claude(openAIResponse *dto.ChatCompletionsStreamRespon return claudeResponses } +func FinalizeStreamResponseOpenAI2Claude(info convmeta.Meta) []*dto.ClaudeResponse { + if info == nil { + info = &convmeta.Values{} + } + state := info.EnsureClaudeConvertInfo() + if state.Done { + return nil + } + + stopReason := stopReasonOpenAI2Claude(state.FinishReason) + if stopReason == "" { + stopReason = "end_turn" + } + responses := stopOpenBlocks(state) + responses = append(responses, + &dto.ClaudeResponse{ + Type: "message_delta", + Usage: buildClaudeUsageFromOpenAIUsage(state.Usage), + Delta: &dto.ClaudeMediaMessage{ + StopReason: kitutil.GetPointer[string](stopReason), + }, + }, + &dto.ClaudeResponse{Type: "message_stop"}, + ) + state.Done = true + return responses +} + func ResponseOpenAI2Claude(openAIResponse *dto.OpenAITextResponse, info convmeta.Meta) *dto.ClaudeResponse { var stopReason string contents := make([]dto.ClaudeMediaMessage, 0) diff --git a/relaykit/relayconvert/response_registry.go b/relaykit/relayconvert/response_registry.go index 66326b36..a2369a61 100644 --- a/relaykit/relayconvert/response_registry.go +++ b/relaykit/relayconvert/response_registry.go @@ -1,15 +1,17 @@ package relayconvert import ( + "context" "errors" "fmt" "reflect" "strings" "sync" - "context" "github.com/QuantumNous/new-api/relaykit/dto" "github.com/QuantumNous/new-api/relaykit/relayconvert/convmeta" + geminichat "github.com/QuantumNous/new-api/relaykit/relayconvert/internal/gemini_chat" + oaichat "github.com/QuantumNous/new-api/relaykit/relayconvert/internal/oai_chat" kitutil "github.com/QuantumNous/new-api/relaykit/relayconvert/kitutil" "github.com/QuantumNous/new-api/relaykit/types" ) @@ -329,6 +331,13 @@ func FinalizeStreamResponse(c context.Context, info convmeta.Meta, state *Respon return nil, nil } + if state.To == types.RelayFormatClaude && info != nil { + claudeInfo := info.EnsureClaudeConvertInfo() + if claudeInfo.Usage == nil { + claudeInfo.Usage = state.Usage() + } + } + values := make([]any, 0) var usage *dto.Usage for i, spec := range state.specs { @@ -474,9 +483,6 @@ func executeStatelessStreamResponseSpec(c context.Context, info convmeta.Meta, f var usage *dto.Usage resultSteps := make([]ResponseStep, 0, len(steps)) for _, step := range steps { - if step.ConvertStreamChunk != nil || step.NewStreamState != nil || step.FinalizeStream != nil { - return nil, fmt.Errorf("response converter %q requires response stream state", step.ID) - } if step.ConvertStream == nil { return nil, fmt.Errorf("response converter %q has no stream implementation", step.ID) } @@ -897,6 +903,15 @@ func convertOAIChatStreamResponseToClaudeMessages(_ context.Context, info convme return StreamResponseOpenAI2Claude(chatResponse, info), canonicalUsageFromResponse(chatResponse), nil } +func finalizeOAIChatStreamResponseToClaudeMessages(_ context.Context, info convmeta.Meta, _ any) ([]any, *dto.Usage, error) { + if info == nil { + info = &convmeta.Values{} + } + usage := info.EnsureClaudeConvertInfo().Usage + responses := oaichat.FinalizeStreamResponseOpenAI2Claude(info) + return streamValuesFromAny(responses), usage, nil +} + func convertOAIChatResponseToGeminiChat(_ context.Context, info convmeta.Meta, response any) (any, *dto.Usage, error) { chatResponse, err := asOAIChatResponse(response) if err != nil { @@ -955,6 +970,41 @@ func convertGeminiChatResponseToOAIChat(_ context.Context, info convmeta.Meta, r return openAIResponse, usage, nil } +func newGeminiChatToOAIChatStreamState(options ResponseStreamOptions) any { + return geminichat.NewGeminiToChatStreamState(options.ID, options.Created) +} + +func convertGeminiChatStreamResponseChunkToOAIChat(_ context.Context, info convmeta.Meta, response any, state any) ([]any, *dto.Usage, error) { + geminiResponse, err := asGeminiChatResponse(response) + if err != nil { + return nil, nil, err + } + streamState, ok := state.(*geminichat.GeminiToChatStreamState) + if !ok || streamState == nil { + return nil, nil, errors.New("Gemini chat to OAI chat stream state is required") + } + usage := UsageFromGeminiMetadata(geminiResponse.GetUsageMetadata(), fallbackPromptTokens(info)) + model := "" + if info != nil && info.HasChannelMeta() { + model = info.GetUpstreamModelName() + } + responses := streamState.ConvertChunk(geminiResponse, model, usage) + return streamValuesFromAny(responses), usage, nil +} + +func finalizeGeminiChatStreamResponseToOAIChat(_ context.Context, info convmeta.Meta, state any) ([]any, *dto.Usage, error) { + streamState, ok := state.(*geminichat.GeminiToChatStreamState) + if !ok || streamState == nil { + return nil, nil, errors.New("Gemini chat to OAI chat stream state is required") + } + model := "" + if info != nil && info.HasChannelMeta() { + model = info.GetUpstreamModelName() + } + responses := streamState.Finalize(model) + return streamValuesFromAny(responses), streamState.Usage(), nil +} + func convertGeminiChatStreamResponseToOAIChat(_ context.Context, info convmeta.Meta, response any) (any, *dto.Usage, error) { geminiResponse, err := asGeminiChatResponse(response) if err != nil { diff --git a/relaykit/relayconvert/terminal_stream_test.go b/relaykit/relayconvert/terminal_stream_test.go new file mode 100644 index 00000000..6ea59260 --- /dev/null +++ b/relaykit/relayconvert/terminal_stream_test.go @@ -0,0 +1,462 @@ +package relayconvert + +import ( + "testing" + + "github.com/QuantumNous/new-api/relaykit/dto" + "github.com/QuantumNous/new-api/relaykit/relayconvert/convmeta" + "github.com/QuantumNous/new-api/relaykit/types" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGeminiToOpenAIStatefulStreamTerminal(t *testing.T) { + tests := []struct { + name string + chunk *dto.GeminiChatResponse + wantFinishReason string + wantFinishOnFinalize bool + wantEmptyFinishDelta bool + }{ + { + name: "stop", + chunk: terminalTestGeminiChunk("Hello", "STOP", false), + wantFinishReason: types.FinishReasonStop, + wantEmptyFinishDelta: true, + }, + { + name: "tool call", + chunk: terminalTestGeminiChunk("", "STOP", true), + wantFinishReason: types.FinishReasonToolCalls, + wantEmptyFinishDelta: true, + }, + { + name: "non stop finish reason", + chunk: terminalTestGeminiChunk("partial", "MAX_TOKENS", false), + wantFinishReason: types.FinishReasonLength, + }, + { + name: "truncated stream", + chunk: terminalTestGeminiChunk("partial", "", false), + wantFinishReason: types.FinishReasonStop, + wantFinishOnFinalize: true, + wantEmptyFinishDelta: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + info := &convmeta.Values{ + ChannelMetaAttached: true, + UpstreamModelName: "upstream-model", + } + state, err := NewResponseStreamState( + types.RelayFormatGemini, + types.RelayFormatOpenAI, + ResponseStreamOptions{ + ID: "chatcmpl-fixed", + Created: 1700000000, + }, + ) + require.NoError(t, err) + + results, err := ConvertStreamResponseChunk(nil, info, state, tt.chunk) + require.NoError(t, err) + chunkFinishes := terminalTestFinishedChatChunks(t, results) + if tt.wantFinishOnFinalize { + assert.Empty(t, chunkFinishes) + } else { + require.Len(t, chunkFinishes, 1) + } + + finalResults, err := FinalizeStreamResponse(nil, info, state) + require.NoError(t, err) + finalFinishes := terminalTestFinishedChatChunks(t, finalResults) + if tt.wantFinishOnFinalize { + require.Len(t, finalFinishes, 1) + } else { + assert.Empty(t, finalFinishes) + } + + finishes := append(chunkFinishes, finalFinishes...) + require.Len(t, finishes, 1) + finish := finishes[0] + require.Len(t, finish.Choices, 1) + require.NotNil(t, finish.Choices[0].FinishReason) + assert.Equal(t, tt.wantFinishReason, *finish.Choices[0].FinishReason) + assert.Equal(t, "chatcmpl-fixed", finish.Id) + assert.Equal(t, int64(1700000000), finish.Created) + assert.Equal(t, "upstream-model", finish.Model) + require.NotNil(t, finish.Usage) + assert.Equal(t, 4, finish.Usage.PromptTokens) + assert.Equal(t, 2, finish.Usage.CompletionTokens) + assert.Equal(t, 6, finish.Usage.TotalTokens) + if tt.wantEmptyFinishDelta { + assert.Nil(t, finish.Choices[0].Delta.Content) + assert.Empty(t, finish.Choices[0].Delta.ToolCalls) + } + + repeatedFinal, err := FinalizeStreamResponse(nil, info, state) + require.NoError(t, err) + assert.Empty(t, repeatedFinal) + }) + } +} + +func TestClaudeTargetStatefulStreamTerminalTail(t *testing.T) { + tests := []struct { + name string + from types.RelayFormat + chunks []any + wantFinalizerTerminals bool + wantStopReason string + }{ + { + name: "gemini to claude", + from: types.RelayFormatGemini, + chunks: []any{ + terminalTestGeminiChunkWithoutUsage("Hello", ""), + terminalTestGeminiChunk(" world", "STOP", false), + }, + wantStopReason: "end_turn", + }, + { + name: "gemini tool call with split usage", + from: types.RelayFormatGemini, + chunks: []any{ + terminalTestGeminiToolChunkWithoutUsage(), + terminalTestGeminiChunk("", "STOP", false), + }, + wantStopReason: "tool_use", + }, + { + name: "gemini non stop finish with split usage", + from: types.RelayFormatGemini, + chunks: []any{ + terminalTestGeminiChunkWithoutUsage("partial", "MAX_TOKENS"), + terminalTestGeminiUsageOnlyChunk(), + }, + wantStopReason: "max_tokens", + }, + { + name: "responses to claude", + from: types.RelayFormatOpenAIResponses, + chunks: []any{ + &dto.ResponsesStreamResponse{ + Type: "response.output_text.delta", + Delta: "Hello", + }, + &dto.ResponsesStreamResponse{ + Type: "response.output_text.delta", + Delta: " world", + }, + &dto.ResponsesStreamResponse{ + Type: "response.completed", + Response: &dto.OpenAIResponsesResponse{ + ID: "resp-fixed", + Object: "response", + Model: "upstream-model", + Status: []byte(`"completed"`), + Usage: &dto.Usage{ + InputTokens: 4, + OutputTokens: 2, + TotalTokens: 6, + }, + }, + }, + }, + wantFinalizerTerminals: true, + wantStopReason: "end_turn", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + info := &convmeta.Values{ + ChannelMetaAttached: true, + UpstreamModelName: "upstream-model", + ClaudeConvertInfo: &convmeta.ClaudeConvertInfo{ + LastMessagesType: convmeta.LastMessageTypeNone, + }, + } + state, err := NewResponseStreamState( + tt.from, + types.RelayFormatClaude, + ResponseStreamOptions{ + ID: "stream-fixed", + Model: "upstream-model", + Created: 1700000000, + }, + ) + require.NoError(t, err) + + var results []ResponseResult + for _, chunk := range tt.chunks { + chunkResults, err := ConvertStreamResponseChunk(nil, info, state, chunk) + require.NoError(t, err) + results = append(results, chunkResults...) + } + + finalResults, err := FinalizeStreamResponse(nil, info, state) + require.NoError(t, err) + if tt.wantFinalizerTerminals { + require.Len(t, finalResults, 3) + } else { + assert.Empty(t, finalResults) + } + results = append(results, finalResults...) + + terminalTestAssertClaudeTail(t, results, tt.wantStopReason) + assert.True(t, info.ClaudeConvertInfo.Done) + + repeatedFinal, err := FinalizeStreamResponse(nil, info, state) + require.NoError(t, err) + assert.Empty(t, repeatedFinal) + }) + } + + t.Run("preserves preseeded usage", func(t *testing.T) { + info := &convmeta.Values{ + ClaudeConvertInfo: &convmeta.ClaudeConvertInfo{ + LastMessagesType: convmeta.LastMessageTypeNone, + }, + } + state, err := NewResponseStreamState( + types.RelayFormatOpenAIResponses, + types.RelayFormatClaude, + ResponseStreamOptions{Model: "upstream-model"}, + ) + require.NoError(t, err) + + chunks := []*dto.ResponsesStreamResponse{ + { + Type: "response.output_text.delta", + Delta: "Hello", + }, + { + Type: "response.completed", + Response: &dto.OpenAIResponsesResponse{ + ID: "resp-fixed", + Object: "response", + Model: "upstream-model", + Status: []byte(`"completed"`), + Usage: &dto.Usage{ + InputTokens: 4, + OutputTokens: 2, + TotalTokens: 6, + }, + }, + }, + } + for _, chunk := range chunks { + _, err := ConvertStreamResponseChunk(nil, info, state, chunk) + require.NoError(t, err) + } + + preseeded := &dto.Usage{ + PromptTokens: 11, + CompletionTokens: 7, + TotalTokens: 18, + } + info.ClaudeConvertInfo.Usage = preseeded + finalResults, err := FinalizeStreamResponse(nil, info, state) + require.NoError(t, err) + require.Len(t, finalResults, 3) + assert.Same(t, preseeded, info.ClaudeConvertInfo.Usage) + + messageDelta, ok := finalResults[1].Value.(*dto.ClaudeResponse) + require.True(t, ok) + assert.Equal(t, "message_delta", messageDelta.Type) + require.NotNil(t, messageDelta.Usage) + assert.Equal(t, 11, messageDelta.Usage.InputTokens) + assert.Equal(t, 7, messageDelta.Usage.OutputTokens) + }) +} + +func TestConvertStreamResponseKeepsStatelessCompatibility(t *testing.T) { + t.Run("gemini to openai", func(t *testing.T) { + info := &convmeta.Values{ + ChannelMetaAttached: true, + UpstreamModelName: "upstream-model", + } + result, err := ConvertStreamResponse( + nil, + info, + types.RelayFormatOpenAI, + terminalTestGeminiChunk("Hello", "STOP", false), + ) + require.NoError(t, err) + require.IsType(t, &dto.ChatCompletionsStreamResponse{}, result.Value) + + response := result.Value.(*dto.ChatCompletionsStreamResponse) + require.Len(t, response.Choices, 1) + require.NotNil(t, response.Choices[0].Delta.Content) + assert.Equal(t, "Hello", *response.Choices[0].Delta.Content) + assert.Nil(t, response.Choices[0].FinishReason) + assert.Equal(t, "upstream-model", response.Model) + require.NotNil(t, response.Usage) + assert.Equal(t, 4, response.Usage.PromptTokens) + assert.Equal(t, 2, response.Usage.CompletionTokens) + assert.Equal(t, 6, response.Usage.TotalTokens) + }) + + t.Run("openai to claude", func(t *testing.T) { + info := &convmeta.Values{ + SendResponseCount: 1, + ClaudeConvertInfo: &convmeta.ClaudeConvertInfo{ + LastMessagesType: convmeta.LastMessageTypeNone, + }, + } + result, err := ConvertStreamResponse( + nil, + info, + types.RelayFormatClaude, + &dto.ChatCompletionsStreamResponse{ + Id: "chatcmpl-fixed", + Object: "chat.completion.chunk", + Created: 1700000000, + Model: "upstream-model", + Choices: []dto.ChatCompletionsStreamResponseChoice{ + { + Delta: dto.ChatCompletionsStreamResponseChoiceDelta{ + Role: "assistant", + Content: terminalTestPtr("Hello"), + }, + }, + }, + Usage: &dto.Usage{ + PromptTokens: 4, + CompletionTokens: 2, + TotalTokens: 6, + }, + }, + ) + require.NoError(t, err) + require.IsType(t, []*dto.ClaudeResponse{}, result.Value) + + responses := result.Value.([]*dto.ClaudeResponse) + require.Len(t, responses, 3) + assert.Equal(t, "message_start", responses[0].Type) + assert.Equal(t, "content_block_start", responses[1].Type) + assert.Equal(t, "content_block_delta", responses[2].Type) + require.NotNil(t, responses[2].Delta) + require.NotNil(t, responses[2].Delta.Text) + assert.Equal(t, "Hello", *responses[2].Delta.Text) + assert.False(t, info.ClaudeConvertInfo.Done) + assert.Equal(t, 6, result.Usage.TotalTokens) + }) +} + +func terminalTestGeminiChunk(text string, finishReason string, toolCall bool) *dto.GeminiChatResponse { + response := terminalTestGeminiChunkWithoutUsage(text, finishReason) + response.HasUsageMetadata = true + response.UsageMetadata = dto.GeminiUsageMetadata{ + PromptTokenCount: 4, + CandidatesTokenCount: 2, + TotalTokenCount: 6, + } + if toolCall { + response.Candidates[0].Content.Parts = []dto.GeminiPart{ + { + FunctionCall: &dto.FunctionCall{ + FunctionName: "lookup", + Arguments: map[string]any{"q": "x"}, + }, + }, + } + } + return response +} + +func terminalTestGeminiChunkWithoutUsage(text string, finishReason string) *dto.GeminiChatResponse { + candidate := dto.GeminiChatCandidate{ + Content: dto.GeminiChatContent{ + Role: "model", + Parts: []dto.GeminiPart{{Text: text}}, + }, + } + if finishReason != "" { + candidate.FinishReason = terminalTestPtr(finishReason) + } + return &dto.GeminiChatResponse{ + Candidates: []dto.GeminiChatCandidate{candidate}, + } +} + +func terminalTestGeminiToolChunkWithoutUsage() *dto.GeminiChatResponse { + return &dto.GeminiChatResponse{ + Candidates: []dto.GeminiChatCandidate{ + { + Content: dto.GeminiChatContent{ + Role: "model", + Parts: []dto.GeminiPart{ + { + FunctionCall: &dto.FunctionCall{ + FunctionName: "lookup", + Arguments: map[string]any{"q": "x"}, + }, + }, + }, + }, + }, + }, + } +} + +func terminalTestGeminiUsageOnlyChunk() *dto.GeminiChatResponse { + return &dto.GeminiChatResponse{ + HasUsageMetadata: true, + UsageMetadata: dto.GeminiUsageMetadata{ + PromptTokenCount: 4, + CandidatesTokenCount: 2, + TotalTokenCount: 6, + }, + } +} + +func terminalTestFinishedChatChunks(t *testing.T, results []ResponseResult) []*dto.ChatCompletionsStreamResponse { + t.Helper() + finished := make([]*dto.ChatCompletionsStreamResponse, 0, 1) + for _, result := range results { + response, ok := result.Value.(*dto.ChatCompletionsStreamResponse) + require.True(t, ok, "unexpected stream result type %T", result.Value) + if response.IsFinished() { + finished = append(finished, response) + } + } + return finished +} + +func terminalTestAssertClaudeTail(t *testing.T, results []ResponseResult, wantStopReason string) { + t.Helper() + responses := make([]*dto.ClaudeResponse, 0, len(results)) + eventCounts := make(map[string]int) + for _, result := range results { + response, ok := result.Value.(*dto.ClaudeResponse) + require.True(t, ok, "unexpected stream result type %T", result.Value) + responses = append(responses, response) + eventCounts[response.Type]++ + } + + require.GreaterOrEqual(t, len(responses), 4) + assert.Equal(t, "message_start", responses[0].Type) + tail := responses[len(responses)-3:] + assert.Equal(t, "content_block_stop", tail[0].Type) + require.NotNil(t, tail[0].Index) + assert.Equal(t, 0, *tail[0].Index) + assert.Equal(t, "message_delta", tail[1].Type) + require.NotNil(t, tail[1].Delta) + require.NotNil(t, tail[1].Delta.StopReason) + assert.Equal(t, wantStopReason, *tail[1].Delta.StopReason) + require.NotNil(t, tail[1].Usage) + assert.Equal(t, 4, tail[1].Usage.InputTokens) + assert.Equal(t, 2, tail[1].Usage.OutputTokens) + assert.Equal(t, "message_stop", tail[2].Type) + assert.Equal(t, 1, eventCounts["content_block_stop"]) + assert.Equal(t, 1, eventCounts["message_delta"]) + assert.Equal(t, 1, eventCounts["message_stop"]) +} + +func terminalTestPtr[T any](value T) *T { + return &value +} diff --git a/relaykit/relayconvert/testdata/golden/stream/gemini_to_claude.golden.json b/relaykit/relayconvert/testdata/golden/stream/gemini_to_claude.golden.json index 9ea6b5dc..2139770c 100644 --- a/relaykit/relayconvert/testdata/golden/stream/gemini_to_claude.golden.json +++ b/relaykit/relayconvert/testdata/golden/stream/gemini_to_claude.golden.json @@ -14,7 +14,7 @@ "claude_cache_creation_1_h_tokens": 0 }, "role": "assistant", - "id": "chatcmpl-", + "id": "stream_fixed", "content": [] } }, @@ -41,6 +41,53 @@ "type": "text_delta", "text": " world" } + }, + { + "type": "content_block_stop", + "index": 0 + }, + { + "type": "message_delta", + "usage": { + "input_tokens": 4, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "output_tokens": 2, + "claude_cache_creation_5_m_tokens": 0, + "claude_cache_creation_1_h_tokens": 0, + "billing_usage": { + "source": "oai_chat", + "semantic": "openai", + "openai_usage": { + "prompt_tokens": 4, + "completion_tokens": 2, + "total_tokens": 6, + "prompt_tokens_details": { + "cached_tokens": 0, + "text_tokens": 4, + "audio_tokens": 0, + "image_tokens": 0 + }, + "completion_tokens_details": { + "text_tokens": 0, + "audio_tokens": 0, + "image_tokens": 0, + "reasoning_tokens": 0 + }, + "input_tokens": 0, + "output_tokens": 0, + "input_tokens_details": null, + "claude_cache_creation_5_m_tokens": 0, + "claude_cache_creation_1_h_tokens": 0 + } + } + }, + "delta": { + "stop_reason": "end_turn" + } + }, + { + "type": "message_stop" } ], "usage": { diff --git a/relaykit/relayconvert/testdata/golden/stream/gemini_to_openai.golden.json b/relaykit/relayconvert/testdata/golden/stream/gemini_to_openai.golden.json index 7794c776..c1718f6b 100644 --- a/relaykit/relayconvert/testdata/golden/stream/gemini_to_openai.golden.json +++ b/relaykit/relayconvert/testdata/golden/stream/gemini_to_openai.golden.json @@ -1,7 +1,7 @@ { "events": [ { - "id": "chatcmpl-", + "id": "stream_fixed", "object": "chat.completion.chunk", "created": 0, "model": "upstream-model", @@ -40,7 +40,7 @@ } }, { - "id": "chatcmpl-", + "id": "stream_fixed", "object": "chat.completion.chunk", "created": 0, "model": "upstream-model", @@ -92,6 +92,58 @@ "claude_cache_creation_5_m_tokens": 0, "claude_cache_creation_1_h_tokens": 0 } + }, + { + "id": "stream_fixed", + "object": "chat.completion.chunk", + "created": 0, + "model": "upstream-model", + "system_fingerprint": null, + "choices": [ + { + "delta": {}, + "logprobs": null, + "finish_reason": "stop", + "index": 0 + } + ], + "usage": { + "prompt_tokens": 4, + "completion_tokens": 2, + "total_tokens": 6, + "billing_usage": { + "source": "gemini_chat", + "semantic": "gemini", + "gemini_usage_metadata": { + "promptTokenCount": 4, + "toolUsePromptTokenCount": 0, + "candidatesTokenCount": 2, + "totalTokenCount": 6, + "thoughtsTokenCount": 0, + "cachedContentTokenCount": 0, + "promptTokensDetails": [], + "toolUsePromptTokensDetails": [], + "candidatesTokensDetails": [] + } + }, + "prompt_tokens_details": { + "cached_tokens": 0, + "text_tokens": 4, + "audio_tokens": 0, + "image_tokens": 0 + }, + "completion_tokens_details": { + "text_tokens": 0, + "audio_tokens": 0, + "image_tokens": 0, + "reasoning_tokens": 0 + }, + "input_tokens": 0, + "output_tokens": 0, + "input_tokens_details": null, + "claude_cache_creation_5_m_tokens": 0, + "claude_cache_creation_1_h_tokens": 0 + } } ], "usage": { diff --git a/relaykit/relayconvert/testdata/golden/stream/openai_responses_to_claude.golden.json b/relaykit/relayconvert/testdata/golden/stream/openai_responses_to_claude.golden.json index 49cd6683..b43ea91e 100644 --- a/relaykit/relayconvert/testdata/golden/stream/openai_responses_to_claude.golden.json +++ b/relaykit/relayconvert/testdata/golden/stream/openai_responses_to_claude.golden.json @@ -41,6 +41,53 @@ "type": "text_delta", "text": " world" } + }, + { + "type": "content_block_stop", + "index": 0 + }, + { + "type": "message_delta", + "usage": { + "input_tokens": 4, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "output_tokens": 2, + "claude_cache_creation_5_m_tokens": 0, + "claude_cache_creation_1_h_tokens": 0, + "billing_usage": { + "source": "oai_responses", + "semantic": "openai", + "openai_usage": { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_tokens": 6, + "prompt_tokens_details": { + "cached_tokens": 0, + "text_tokens": 0, + "audio_tokens": 0, + "image_tokens": 0 + }, + "completion_tokens_details": { + "text_tokens": 0, + "audio_tokens": 0, + "image_tokens": 0, + "reasoning_tokens": 0 + }, + "input_tokens": 4, + "output_tokens": 2, + "input_tokens_details": null, + "claude_cache_creation_5_m_tokens": 0, + "claude_cache_creation_1_h_tokens": 0 + } + } + }, + "delta": { + "stop_reason": "end_turn" + } + }, + { + "type": "message_stop" } ], "usage": { diff --git a/relaykit/relayconvert/text_converter_registry.go b/relaykit/relayconvert/text_converter_registry.go index 49e19379..dedbd2e3 100644 --- a/relaykit/relayconvert/text_converter_registry.go +++ b/relaykit/relayconvert/text_converter_registry.go @@ -70,9 +70,10 @@ var builtinTextConverters = []TextConverterSpec{ Convert: convertOpenAIRequestToClaude, }, Resp: TextResponseSide{ - Convert: convertOAIChatResponseToClaudeMessages, - ConvertStream: convertOAIChatStreamResponseToClaudeMessages, - Aliases: []string{ResponseConverterOAIChatToClaudeMessages}, + Convert: convertOAIChatResponseToClaudeMessages, + ConvertStream: convertOAIChatStreamResponseToClaudeMessages, + FinalizeStream: finalizeOAIChatStreamResponseToClaudeMessages, + Aliases: []string{ResponseConverterOAIChatToClaudeMessages}, }, }, { @@ -84,9 +85,12 @@ var builtinTextConverters = []TextConverterSpec{ Convert: convertGeminiRequestToOpenAI, }, Resp: TextResponseSide{ - Convert: convertGeminiChatResponseToOAIChat, - ConvertStream: convertGeminiChatStreamResponseToOAIChat, - Aliases: []string{ResponseConverterGeminiChatToOAIChat}, + Convert: convertGeminiChatResponseToOAIChat, + ConvertStream: convertGeminiChatStreamResponseToOAIChat, + NewStreamState: newGeminiChatToOAIChatStreamState, + ConvertStreamChunk: convertGeminiChatStreamResponseChunkToOAIChat, + FinalizeStream: finalizeGeminiChatStreamResponseToOAIChat, + Aliases: []string{ResponseConverterGeminiChatToOAIChat}, }, }, { diff --git a/relaykit/relayconvert/text_converter_registry_test.go b/relaykit/relayconvert/text_converter_registry_test.go index 53979572..2f569085 100644 --- a/relaykit/relayconvert/text_converter_registry_test.go +++ b/relaykit/relayconvert/text_converter_registry_test.go @@ -23,7 +23,7 @@ func TestLookupBuiltinTextConverters(t *testing.T) { }{ {id: ConverterClaudeMessagesToOpenAIChat, from: types.RelayFormatClaude, to: types.RelayFormatOpenAI, quality: TextConverterQualityFair, reqDirect: true, respDirect: true, respAlias: ResponseConverterClaudeMessagesToOAIChat}, {id: ConverterOpenAIChatToClaudeMessages, from: types.RelayFormatOpenAI, to: types.RelayFormatClaude, quality: TextConverterQualityFair, reqDirect: true, respDirect: true, respAlias: ResponseConverterOAIChatToClaudeMessages}, - {id: ConverterGeminiContentToOpenAIChat, from: types.RelayFormatGemini, to: types.RelayFormatOpenAI, quality: TextConverterQualityFair, reqDirect: true, respDirect: true, respAlias: ResponseConverterGeminiChatToOAIChat}, + {id: ConverterGeminiContentToOpenAIChat, from: types.RelayFormatGemini, to: types.RelayFormatOpenAI, quality: TextConverterQualityFair, reqDirect: true, respDirect: true, respAlias: ResponseConverterGeminiChatToOAIChat, streamDirect: true}, {id: ConverterOpenAIChatToGeminiContent, from: types.RelayFormatOpenAI, to: types.RelayFormatGemini, quality: TextConverterQualityFair, reqDirect: true, respDirect: true, respAlias: ResponseConverterOAIChatToGeminiChat}, {id: ConverterOpenAIChatToOpenAIResponses, from: types.RelayFormatOpenAI, to: types.RelayFormatOpenAIResponses, quality: TextConverterQualityGood, reqDirect: true, respDirect: true, respAlias: ResponseConverterOAIChatToOAIResponses, streamDirect: true}, {id: ConverterOpenAIResponsesToOpenAIChat, from: types.RelayFormatOpenAIResponses, to: types.RelayFormatOpenAI, quality: TextConverterQualityGood, reqDirect: true, respDirect: true, respAlias: ResponseConverterOAIResponsesToOAIChat, streamDirect: true},