diff --git a/core/relay/adaptor/openai/chat.go b/core/relay/adaptor/openai/chat.go index 129df91d..b64c7f97 100644 --- a/core/relay/adaptor/openai/chat.go +++ b/core/relay/adaptor/openai/chat.go @@ -23,13 +23,73 @@ import ( // chatCompletionStreamState manages state for ChatCompletion stream conversion type chatCompletionStreamState struct { - messageID string - meta *meta.Meta - c *gin.Context - currentToolCall *relaymodel.ToolCall - currentToolCallID string - toolCallArgs string - hasToolCall bool + messageID string + created int64 + meta *meta.Meta + c *gin.Context + toolCallIndexByItemID map[string]int + toolCallIndexByOutputIndex map[int]int + nextToolCallIndex int + hasToolCall bool +} + +func (s *chatCompletionStreamState) createdAt() int64 { + if s.created == 0 { + s.created = time.Now().Unix() + } + + return s.created +} + +func (s *chatCompletionStreamState) registerToolCall( + event *relaymodel.ResponseStreamEvent, +) int { + if event.Item.ID != "" { + if index, ok := s.toolCallIndexByItemID[event.Item.ID]; ok { + return index + } + } + + if event.OutputIndex != nil { + if index, ok := s.toolCallIndexByOutputIndex[*event.OutputIndex]; ok { + return index + } + } + + index := s.nextToolCallIndex + s.nextToolCallIndex++ + + if event.Item.ID != "" { + s.toolCallIndexByItemID[event.Item.ID] = index + } + + if event.OutputIndex != nil { + s.toolCallIndexByOutputIndex[*event.OutputIndex] = index + } + + return index +} + +func (s *chatCompletionStreamState) toolCallIndex( + event *relaymodel.ResponseStreamEvent, +) (int, bool) { + if event.ItemID != "" { + index, ok := s.toolCallIndexByItemID[event.ItemID] + if ok { + return index, true + } + } + + if event.OutputIndex != nil { + index, ok := s.toolCallIndexByOutputIndex[*event.OutputIndex] + return index, ok + } + + if event.ItemID == "" && s.nextToolCallIndex > 0 { + return s.nextToolCallIndex - 1, true + } + + return 0, false } func responseModelName(meta *meta.Meta) string { @@ -132,11 +192,14 @@ func (s *chatCompletionStreamState) handleResponseCreated( } s.messageID = event.Response.ID + if event.Response.CreatedAt != 0 { + s.created = event.Response.CreatedAt + } return &relaymodel.ChatCompletionsStreamResponse{ ID: s.messageID, Object: relaymodel.ChatCompletionChunkObject, - Created: event.Response.CreatedAt, + Created: s.createdAt(), Model: responseModelName(s.meta), Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{ { @@ -160,7 +223,7 @@ func (s *chatCompletionStreamState) handleOutputTextDelta( return &relaymodel.ChatCompletionsStreamResponse{ ID: s.messageID, Object: relaymodel.ChatCompletionChunkObject, - Created: time.Now().Unix(), + Created: s.createdAt(), Model: responseModelName(s.meta), Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{ { @@ -183,7 +246,7 @@ func (s *chatCompletionStreamState) handleReasoningSummaryTextDelta( return &relaymodel.ChatCompletionsStreamResponse{ ID: s.messageID, Object: relaymodel.ChatCompletionChunkObject, - Created: time.Now().Unix(), + Created: s.createdAt(), Model: responseModelName(s.meta), Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{ { @@ -207,22 +270,13 @@ func (s *chatCompletionStreamState) handleOutputItemAdded( // Track function calls if event.Item.Type == relaymodel.InputItemTypeFunctionCall { s.hasToolCall = true - s.currentToolCallID = event.Item.ID - s.currentToolCall = &relaymodel.ToolCall{ - ID: event.Item.CallID, - Type: relaymodel.ToolChoiceTypeFunction, - Function: relaymodel.Function{ - Name: event.Item.Name, - Arguments: "", - }, - } - s.toolCallArgs = "" + toolCallIndex := s.registerToolCall(event) // Send tool call start return &relaymodel.ChatCompletionsStreamResponse{ ID: s.messageID, Object: relaymodel.ChatCompletionChunkObject, - Created: time.Now().Unix(), + Created: s.createdAt(), Model: responseModelName(s.meta), Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{ { @@ -230,7 +284,7 @@ func (s *chatCompletionStreamState) handleOutputItemAdded( Delta: relaymodel.Message{ ToolCalls: []relaymodel.ToolCall{ { - Index: 0, + Index: toolCallIndex, ID: event.Item.CallID, Type: relaymodel.ToolChoiceTypeFunction, Function: relaymodel.Function{ @@ -249,7 +303,7 @@ func (s *chatCompletionStreamState) handleOutputItemAdded( return &relaymodel.ChatCompletionsStreamResponse{ ID: s.messageID, Object: relaymodel.ChatCompletionChunkObject, - Created: time.Now().Unix(), + Created: s.createdAt(), Model: responseModelName(s.meta), Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{ { @@ -269,18 +323,20 @@ func (s *chatCompletionStreamState) handleOutputItemAdded( func (s *chatCompletionStreamState) handleFunctionCallArgumentsDelta( event *relaymodel.ResponseStreamEvent, ) *relaymodel.ChatCompletionsStreamResponse { - if event.Delta == "" || s.currentToolCall == nil { + if event.Delta == "" { return nil } - // Accumulate arguments - s.toolCallArgs += event.Delta + toolCallIndex, ok := s.toolCallIndex(event) + if !ok { + return nil + } // Send delta return &relaymodel.ChatCompletionsStreamResponse{ ID: s.messageID, Object: relaymodel.ChatCompletionChunkObject, - Created: time.Now().Unix(), + Created: s.createdAt(), Model: responseModelName(s.meta), Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{ { @@ -288,7 +344,7 @@ func (s *chatCompletionStreamState) handleFunctionCallArgumentsDelta( Delta: relaymodel.Message{ ToolCalls: []relaymodel.ToolCall{ { - Index: 0, + Index: toolCallIndex, Function: relaymodel.Function{ Arguments: event.Delta, }, @@ -300,32 +356,6 @@ func (s *chatCompletionStreamState) handleFunctionCallArgumentsDelta( } } -// handleOutputItemDone handles response.output_item.done event for ChatCompletion -func (s *chatCompletionStreamState) handleOutputItemDone( - event *relaymodel.ResponseStreamEvent, -) { - if event.Item == nil { - return - } - - // Handle function call completion - if event.Item.Type == relaymodel.InputItemTypeFunctionCall && s.currentToolCall != nil && - event.Item.ID == s.currentToolCallID { - // Update with final arguments - if s.toolCallArgs != "" { - s.currentToolCall.Function.Arguments = s.toolCallArgs - } - - // Reset state - s.currentToolCall = nil - s.currentToolCallID = "" - s.toolCallArgs = "" - - // No need to send another chunk - arguments already streamed - return - } -} - // handleResponseCompleted handles response.completed/done event for ChatCompletion func (s *chatCompletionStreamState) handleResponseCompleted( event *relaymodel.ResponseStreamEvent, @@ -344,7 +374,7 @@ func (s *chatCompletionStreamState) handleResponseCompleted( return &relaymodel.ChatCompletionsStreamResponse{ ID: s.messageID, Object: relaymodel.ChatCompletionChunkObject, - Created: time.Now().Unix(), + Created: s.createdAt(), Model: responseModelName(s.meta), Choices: []*relaymodel.ChatCompletionsStreamResponseChoice{ { @@ -1341,6 +1371,8 @@ func ConvertResponsesToChatCompletionResponse( reasonContent := responseReasoningSummaryText(&responsesResp) + var toolCallChoice *relaymodel.TextResponseChoice + // Convert output items to choices for _, outputItem := range responsesResp.Output { switch outputItem.Type { @@ -1380,31 +1412,36 @@ func ConvertResponsesToChatCompletionResponse( toolCallID = outputItem.ID } - finishReason := responseToChatFinishReason(&responsesResp) - if finishReason == relaymodel.FinishReasonStop { - finishReason = relaymodel.FinishReasonToolCalls + if toolCallChoice == nil { + finishReason := responseToChatFinishReason(&responsesResp) + if finishReason == relaymodel.FinishReasonStop { + finishReason = relaymodel.FinishReasonToolCalls + } + + toolCallChoice = &relaymodel.TextResponseChoice{ + Index: len(chatResp.Choices), + Message: relaymodel.Message{ + Role: relaymodel.RoleAssistant, + ReasoningContent: reasonContent, + }, + FinishReason: finishReason, + } + chatResp.Choices = append(chatResp.Choices, toolCallChoice) + reasonContent = "" } - chatResp.Choices = append(chatResp.Choices, &relaymodel.TextResponseChoice{ - Index: len(chatResp.Choices), - Message: relaymodel.Message{ - Role: relaymodel.RoleAssistant, - ReasoningContent: reasonContent, - ToolCalls: []relaymodel.ToolCall{ - { - Index: 0, - ID: toolCallID, - Type: relaymodel.ToolChoiceTypeFunction, - Function: relaymodel.Function{ - Name: outputItem.Name, - Arguments: outputItem.Arguments.String(), - }, - }, + toolCallChoice.Message.ToolCalls = append( + toolCallChoice.Message.ToolCalls, + relaymodel.ToolCall{ + Index: len(toolCallChoice.Message.ToolCalls), + ID: toolCallID, + Type: relaymodel.ToolChoiceTypeFunction, + Function: relaymodel.Function{ + Name: outputItem.Name, + Arguments: outputItem.Arguments.String(), }, }, - FinishReason: finishReason, - }) - reasonContent = "" + ) default: continue @@ -1556,8 +1593,10 @@ func ConvertResponsesToChatCompletionStreamResponse( errorState := responsesStreamErrorState{} state := &chatCompletionStreamState{ - meta: meta, - c: c, + meta: meta, + c: c, + toolCallIndexByItemID: make(map[string]int), + toolCallIndexByOutputIndex: make(map[int]int), } stopStream := false @@ -1635,8 +1674,6 @@ func ConvertResponsesToChatCompletionStreamResponse( chatStreamResp = state.handleOutputItemAdded(&event) case relaymodel.EventFunctionCallArgumentsDelta: chatStreamResp = state.handleFunctionCallArgumentsDelta(&event) - case relaymodel.EventOutputItemDone: - state.handleOutputItemDone(&event) case relaymodel.EventResponseCompleted, relaymodel.EventResponseIncomplete, relaymodel.EventResponseDone: diff --git a/core/relay/adaptor/openai/chat_test.go b/core/relay/adaptor/openai/chat_test.go index 23cadd81..c9feeb8f 100644 --- a/core/relay/adaptor/openai/chat_test.go +++ b/core/relay/adaptor/openai/chat_test.go @@ -1016,6 +1016,51 @@ func TestConvertResponsesToChatCompletionResponse(t *testing.T) { }, expectedStatus: http.StatusOK, }, + { + name: "multiple function calls share one choice", + responsesResp: relaymodel.Response{ + ID: "resp_tools", + Model: "gpt-5-mini", + Status: relaymodel.ResponseStatusCompleted, + CreatedAt: 1781355958, + Output: []relaymodel.OutputItem{ + { + ID: "fc_logs", + Type: relaymodel.InputItemTypeFunctionCall, + CallID: "call_logs", + Name: "command", + Arguments: `{"command":"kubectl logs ..."}`, + }, + { + ID: "fc_statefulset", + Type: relaymodel.InputItemTypeFunctionCall, + CallID: "call_statefulset", + Name: "command", + Arguments: `{"command":"kubectl get statefulset ..."}`, + }, + }, + Usage: &relaymodel.ResponseUsage{ + InputTokens: 12, + OutputTokens: 6, + TotalTokens: 18, + }, + }, + checkFunc: func(t *testing.T, chatResp relaymodel.TextResponse) { + t.Helper() + require.Len(t, chatResp.Choices, 1) + + choice := chatResp.Choices[0] + assert.Equal(t, relaymodel.FinishReasonToolCalls, choice.FinishReason) + require.Len(t, choice.Message.ToolCalls, 2) + assert.Equal(t, 0, choice.Message.ToolCalls[0].Index) + assert.Equal(t, `{"command":"kubectl logs ..."}`, + choice.Message.ToolCalls[0].Function.Arguments) + assert.Equal(t, 1, choice.Message.ToolCalls[1].Index) + assert.Equal(t, `{"command":"kubectl get statefulset ..."}`, + choice.Message.ToolCalls[1].Function.Arguments) + }, + expectedStatus: http.StatusOK, + }, { name: "incomplete function call keeps incomplete finish reason", responsesResp: relaymodel.Response{ @@ -1583,6 +1628,89 @@ func TestConvertResponsesToChatCompletionStreamResponseUsesToolCallsFinishReason assert.Equal(t, 1, strings.Count(w.Body.String(), "data: [DONE]")) } +func TestConvertResponsesToChatCompletionStreamResponseKeepsParallelToolCallsSeparate( + t *testing.T, +) { + gin.SetMode(gin.TestMode) + + stream := strings.Join([]string{ + `event: response.created`, + `data: {"type":"response.created","response":{"id":"resp_tools","object":"response","created_at":1781355623,"status":"in_progress","model":"gpt-5-mini","output":[],"parallel_tool_calls":true,"store":false}}`, + "", + `event: response.output_item.added`, + `data: {"type":"response.output_item.added","item":{"id":"fc_logs","type":"function_call","call_id":"call_logs","name":"command","arguments":"","status":"in_progress"},"output_index":0,"sequence_number":1}`, + "", + `event: response.output_item.added`, + `data: {"type":"response.output_item.added","item":{"id":"fc_statefulset","type":"function_call","call_id":"call_statefulset","name":"command","arguments":"","status":"in_progress"},"output_index":1,"sequence_number":2}`, + "", + `event: response.function_call_arguments.delta`, + `data: {"type":"response.function_call_arguments.delta","item_id":"fc_statefulset","output_index":1,"delta":"{\"command\":\"kubectl get statefulset ...\"}","sequence_number":3}`, + "", + `event: response.function_call_arguments.delta`, + `data: {"type":"response.function_call_arguments.delta","item_id":"fc_logs","output_index":0,"delta":"{\"command\":\"kubectl logs ...\"}","sequence_number":4}`, + "", + `event: response.output_item.done`, + `data: {"type":"response.output_item.done","item":{"id":"fc_logs","type":"function_call","call_id":"call_logs","name":"command","arguments":"{\"command\":\"kubectl logs ...\"}","status":"completed"},"output_index":0,"sequence_number":5}`, + "", + `event: response.output_item.done`, + `data: {"type":"response.output_item.done","item":{"id":"fc_statefulset","type":"function_call","call_id":"call_statefulset","name":"command","arguments":"{\"command\":\"kubectl get statefulset ...\"}","status":"completed"},"output_index":1,"sequence_number":6}`, + "", + `event: response.completed`, + `data: {"type":"response.completed","response":{"id":"resp_tools","object":"response","created_at":1781355623,"status":"completed","model":"gpt-5-mini","output":[],"parallel_tool_calls":true,"store":false,"usage":{"input_tokens":12,"output_tokens":6,"total_tokens":18}},"sequence_number":7}`, + "", + }, "\n") + + httpResp := &http.Response{ + StatusCode: http.StatusOK, + Body: &mockReadCloser{Reader: bytes.NewReader([]byte(stream))}, + Header: make(http.Header), + } + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequestWithContext( + t.Context(), + http.MethodPost, + "/v1/chat/completions", + nil, + ) + + _, err := openai.ConvertResponsesToChatCompletionStreamResponse( + &meta.Meta{ActualModel: "gpt-5-mini"}, + c, + httpResp, + ) + require.Nil(t, err) + + chunks := collectChatCompletionStreamChunks(t, w.Body.String()) + require.Len(t, chunks, 6) + + for _, chunk := range chunks { + assert.Equal(t, int64(1781355623), chunk.Created) + } + + assert.NotContains(t, w.Body.String(), `"id":""`) + assert.NotContains(t, w.Body.String(), `"type":""`) + + argumentsByIndex := make(map[int]string) + + callIDsByIndex := make(map[int]string) + for _, chunk := range chunks { + for _, toolCall := range chunk.Choices[0].Delta.ToolCalls { + argumentsByIndex[toolCall.Index] += toolCall.Function.Arguments + if toolCall.ID != "" { + callIDsByIndex[toolCall.Index] = toolCall.ID + } + } + } + + assert.Equal(t, "call_logs", callIDsByIndex[0]) + assert.Equal(t, `{"command":"kubectl logs ..."}`, argumentsByIndex[0]) + assert.Equal(t, "call_statefulset", callIDsByIndex[1]) + assert.Equal(t, `{"command":"kubectl get statefulset ..."}`, argumentsByIndex[1]) + assert.Equal(t, relaymodel.FinishReasonToolCalls, chunks[5].Choices[0].FinishReason) +} + func TestConvertResponsesToChatCompletionStreamResponseUsesOriginModelForEveryChunk( t *testing.T, ) { diff --git a/core/relay/model/chat_test.go b/core/relay/model/chat_test.go index 8c910c3c..b7145e68 100644 --- a/core/relay/model/chat_test.go +++ b/core/relay/model/chat_test.go @@ -10,6 +10,23 @@ import ( "github.com/smartystreets/goconvey/convey" ) +func TestToolCallOmitsEmptyDeltaFields(t *testing.T) { + data, err := json.Marshal(model.ToolCall{ + Index: 0, + Function: model.Function{ + Arguments: `{"command":"kubectl logs ..."}`, + }, + }) + if err != nil { + t.Fatal(err) + } + + expected := `{"index":0,"function":{"arguments":"{\"command\":\"kubectl logs ...\"}"}}` + if string(data) != expected { + t.Fatalf("unexpected tool call JSON: %s", data) + } +} + func TestChatUsage(t *testing.T) { convey.Convey("ChatUsage", t, func() { convey.Convey("ToModelUsage", func() { diff --git a/core/relay/model/tool.go b/core/relay/model/tool.go index 3666817c..d1c21650 100644 --- a/core/relay/model/tool.go +++ b/core/relay/model/tool.go @@ -22,8 +22,8 @@ type ExtraContent struct { type ToolCall struct { Index int `json:"index"` - ID string `json:"id"` - Type string `json:"type"` + ID string `json:"id,omitempty"` + Type string `json:"type,omitempty"` Function Function `json:"function"` ExtraContent *ExtraContent `json:"extra_content,omitempty"` }