diff --git a/internal/runtime/executor/kimi_executor_test.go b/internal/runtime/executor/kimi_executor_test.go index dea518d4c..c59fe1f33 100644 --- a/internal/runtime/executor/kimi_executor_test.go +++ b/internal/runtime/executor/kimi_executor_test.go @@ -73,6 +73,68 @@ func TestKimiExecutorClaudeRequestPreservesInternalModelSemantics(t *testing.T) } } +func TestKimiExecutorPreservesAssistantContentAndToolCallsFromResponsesHistory(t *testing.T) { + var upstreamBody []byte + ctx := context.WithValue(context.Background(), "cliproxy.roundtripper", kimiRoundTripperFunc(func(req *http.Request) (*http.Response, error) { + var errRead error + upstreamBody, errRead = io.ReadAll(req.Body) + if errRead != nil { + return nil, errRead + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader( + `{"id":"chatcmpl_test","object":"chat.completion","created":1,"model":"k3","choices":[{"index":0,"message":{"role":"assistant","content":"done"},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`, + )), + }, nil + })) + + executor := NewKimiExecutor(&config.Config{}) + auth := &cliproxyauth.Auth{ + Attributes: map[string]string{}, + Metadata: map[string]any{"access_token": "test-token"}, + } + payload := []byte(`{ + "model":"kimi-k3", + "input":[ + {"type":"reasoning","id":"rs_1","summary":[{"type":"summary_text","text":"inspect the next step"}]}, + {"type":"message","role":"assistant","content":[{"type":"output_text","text":"Step 3 completed; continue to step 4."}]}, + {"type":"function_call","call_id":"call_4","name":"exec_command","arguments":"{\"cmd\":\"pwd\"}"}, + {"type":"function_call_output","call_id":"call_4","output":"ok"} + ] + }`) + + _, err := executor.Execute(ctx, auth, cliproxyexecutor.Request{ + Model: "kimi-k3", + Payload: payload, + }, cliproxyexecutor.Options{ + SourceFormat: sdktranslator.FormatOpenAIResponse, + OriginalRequest: payload, + }) + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + + messages := gjson.GetBytes(upstreamBody, "messages").Array() + if got := len(messages); got != 2 { + t.Fatalf("upstream messages count = %d, want 2; body=%s", got, upstreamBody) + } + assistant := messages[0] + if got := assistant.Get("content.0.text").String(); got != "Step 3 completed; continue to step 4." { + t.Fatalf("assistant content = %q, want preserved text; body=%s", got, upstreamBody) + } + if got := assistant.Get("reasoning_content").String(); got != "inspect the next step" { + t.Fatalf("assistant reasoning_content = %q, want inspect the next step; body=%s", got, upstreamBody) + } + if got := assistant.Get("tool_calls.0.id").String(); got != "call_4" { + t.Fatalf("assistant tool call ID = %q, want call_4; body=%s", got, upstreamBody) + } + if got := messages[1].Get("tool_call_id").String(); got != "call_4" { + t.Fatalf("tool output call ID = %q, want call_4; body=%s", got, upstreamBody) + } +} + func TestKimiExecutorCountTokensUsesCanonicalUpstreamModel(t *testing.T) { var upstreamRequest *http.Request var upstreamBody []byte diff --git a/internal/translator/openai/openai/responses/openai_openai-responses_request.go b/internal/translator/openai/openai/responses/openai_openai-responses_request.go index 7ab2ced99..26757b9b6 100644 --- a/internal/translator/openai/openai/responses/openai_openai-responses_request.go +++ b/internal/translator/openai/openai/responses/openai_openai-responses_request.go @@ -85,6 +85,7 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu pendingReasoningContent := "" awaitingToolOutputs := make(map[string]struct{}) deferredMessages := make([][]byte, 0) + mergeableAssistantIndex := -1 takePendingReasoningContent := func() string { reasoningContent := pendingReasoningContent @@ -95,12 +96,29 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu if len(pendingToolCalls) == 0 { return } - assistantMessage := []byte(`{"role":"assistant","tool_calls":[]}`) - assistantMessage, _ = sjson.SetBytes(assistantMessage, "tool_calls", pendingToolCalls) - if reasoningContent := takePendingReasoningContent(); reasoningContent != "" { - assistantMessage, _ = sjson.SetBytes(assistantMessage, "reasoning_content", reasoningContent) + + reasoningContent := takePendingReasoningContent() + mergedIntoAssistant := false + if mergeableAssistantIndex >= 0 && mergeableAssistantIndex == len(messages)-1 { + assistantMessage := gjson.ParseBytes(messages[mergeableAssistantIndex]) + if assistantMessage.Get("role").String() == "assistant" && !assistantMessage.Get("tool_calls").Exists() { + updatedMessage, _ := sjson.SetBytes(messages[mergeableAssistantIndex], "tool_calls", pendingToolCalls) + combinedReasoning := combineOpenAIResponsesReasoning(assistantMessage.Get("reasoning_content").String(), reasoningContent) + if combinedReasoning != "" { + updatedMessage, _ = sjson.SetBytes(updatedMessage, "reasoning_content", combinedReasoning) + } + messages[mergeableAssistantIndex] = updatedMessage + mergedIntoAssistant = true + } + } + if !mergedIntoAssistant { + assistantMessage := []byte(`{"role":"assistant","tool_calls":[]}`) + assistantMessage, _ = sjson.SetBytes(assistantMessage, "tool_calls", pendingToolCalls) + if reasoningContent != "" { + assistantMessage, _ = sjson.SetBytes(assistantMessage, "reasoning_content", reasoningContent) + } + appendMessage(assistantMessage) } - appendMessage(assistantMessage) for _, id := range pendingToolCallIDs { if strings.TrimSpace(id) == "" { continue @@ -109,6 +127,7 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu } pendingToolCalls = pendingToolCalls[:0] pendingToolCallIDs = pendingToolCallIDs[:0] + mergeableAssistantIndex = -1 } flushDeferredMessages := func() { for _, message := range deferredMessages { @@ -124,14 +143,15 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu } return false } - appendRegularMessage := func(message []byte) { + appendRegularMessage := func(message []byte) int { // Keep tool-call adjacency strict for providers that require // assistant(tool_calls) -> tool(tool_call_id) with no message in between. if hasAwaitingToolOutput() { deferredMessages = append(deferredMessages, message) - return + return -1 } appendMessage(message) + return len(messages) - 1 } appendPendingReasoningMessage := func() { reasoningContent := takePendingReasoningContent() @@ -159,6 +179,7 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu if role == "developer" { role = "user" } + mergeableAssistantIndex = -1 if role != "assistant" { appendPendingReasoningMessage() } @@ -196,28 +217,23 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu } if role == "assistant" { - reasoningContent := item.Get("reasoning_content").String() - if reasoningContent == "" { - reasoningContent = takePendingReasoningContent() - } else { - pendingReasoningContent = "" - } + reasoningContent := combineOpenAIResponsesReasoning(takePendingReasoningContent(), item.Get("reasoning_content").String()) if reasoningContent != "" { message, _ = sjson.SetBytes(message, "reasoning_content", reasoningContent) } } - appendRegularMessage(message) + messageIndex := appendRegularMessage(message) + if role == "assistant" { + mergeableAssistantIndex = messageIndex + } case "reasoning": reasoningContent := collectOpenAIResponsesReasoningContent(item) - if pendingReasoningContent == "" { - pendingReasoningContent = reasoningContent - } else { - pendingReasoningContent += reasoningContent - } + pendingReasoningContent = combineOpenAIResponsesReasoning(pendingReasoningContent, reasoningContent) case "function_call": + pendingReasoningContent = combineOpenAIResponsesReasoning(pendingReasoningContent, item.Get("reasoning_content").String()) // Buffer consecutive function calls and emit them as one assistant message. toolCall := []byte(`{"id":"","type":"function","function":{"name":"","arguments":""}}`) @@ -242,6 +258,7 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu } case "function_call_output": + mergeableAssistantIndex = -1 // Handle function call output conversion to tool message toolMessage := []byte(`{"role":"tool","tool_call_id":"","content":""}`) callID := "" @@ -264,6 +281,7 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu } case "custom_tool_call": + pendingReasoningContent = combineOpenAIResponsesReasoning(pendingReasoningContent, item.Get("reasoning_content").String()) // Codex freeform tool call replay: wrap the raw input so it // matches the {"input": string} function shape used when // converting custom tool definitions. @@ -278,6 +296,7 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu } case "custom_tool_call_output": + mergeableAssistantIndex = -1 toolMessage := []byte(`{"role":"tool","tool_call_id":"","content":""}`) callID := strings.TrimSpace(item.Get("call_id").String()) toolMessage, _ = sjson.SetBytes(toolMessage, "tool_call_id", callID) @@ -289,6 +308,9 @@ func ConvertOpenAIResponsesRequestToOpenAIChatCompletions(modelName string, inpu if len(awaitingToolOutputs) == 0 && len(deferredMessages) > 0 { flushDeferredMessages() } + + default: + mergeableAssistantIndex = -1 } } @@ -503,3 +525,21 @@ func collectOpenAIResponsesReasoningContent(item gjson.Result) string { } return reasoningText.String() } + +func combineOpenAIResponsesReasoning(existing, incoming string) string { + existingTrimmed := strings.TrimSpace(existing) + incomingTrimmed := strings.TrimSpace(incoming) + + switch { + case existingTrimmed == "": + return incoming + case incomingTrimmed == "": + return existing + case existingTrimmed == "[reasoning unavailable]": + return incoming + case incomingTrimmed == "[reasoning unavailable]", existingTrimmed == incomingTrimmed: + return existing + default: + return existing + "\n\n" + incoming + } +} diff --git a/internal/translator/openai/openai/responses/openai_openai-responses_request_test.go b/internal/translator/openai/openai/responses/openai_openai-responses_request_test.go index 48fdcc91d..880f92708 100644 --- a/internal/translator/openai/openai/responses/openai_openai-responses_request_test.go +++ b/internal/translator/openai/openai/responses/openai_openai-responses_request_test.go @@ -299,6 +299,119 @@ func TestConvertOpenAIResponsesRequestToOpenAIChatCompletions_AttachesReasoningT } } +func TestConvertOpenAIResponsesRequestToOpenAIChatCompletions_PreservesAssistantContentWithToolCalls(t *testing.T) { + raw := []byte(`{ + "input": [ + { + "type": "reasoning", + "id": "rs_1", + "summary": [{"type": "summary_text", "text": "inspect the next step"}] + }, + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "Step 3 completed; continue to step 4."}] + }, + {"type":"function_call","call_id":"call_4","name":"exec_command","arguments":"{\"cmd\":\"pwd\"}"}, + {"type":"function_call_output","call_id":"call_4","output":"ok"} + ] + }`) + + out := ConvertOpenAIResponsesRequestToOpenAIChatCompletions("kimi-k3", raw, false) + + messages := gjson.GetBytes(out, "messages").Array() + if got := len(messages); got != 2 { + t.Fatalf("messages count = %d, want 2; output=%s", got, out) + } + assistant := messages[0] + if got := assistant.Get("role").String(); got != "assistant" { + t.Fatalf("assistant role = %q, want assistant; output=%s", got, out) + } + if got := assistant.Get("reasoning_content").String(); got != "inspect the next step" { + t.Fatalf("assistant reasoning_content = %q, want inspect the next step; output=%s", got, out) + } + if got := assistant.Get("content.0.text").String(); got != "Step 3 completed; continue to step 4." { + t.Fatalf("assistant content = %q, want preserved text; output=%s", got, out) + } + if got := assistant.Get("tool_calls.0.id").String(); got != "call_4" { + t.Fatalf("assistant tool call ID = %q, want call_4; output=%s", got, out) + } + if got := messages[1].Get("tool_call_id").String(); got != "call_4" { + t.Fatalf("tool output call ID = %q, want call_4; output=%s", got, out) + } +} + +func TestConvertOpenAIResponsesRequestToOpenAIChatCompletions_DoesNotMergeToolCallsAcrossUserMessage(t *testing.T) { + raw := []byte(`{ + "input": [ + {"type":"message","role":"assistant","content":[{"type":"output_text","text":"done"}]}, + {"type":"message","role":"user","content":[{"type":"input_text","text":"next"}]}, + {"type":"function_call","call_id":"call_next","name":"exec_command","arguments":"{}"}, + {"type":"function_call_output","call_id":"call_next","output":"ok"} + ] + }`) + + out := ConvertOpenAIResponsesRequestToOpenAIChatCompletions("kimi-k3", raw, false) + + messages := gjson.GetBytes(out, "messages").Array() + if got := len(messages); got != 4 { + t.Fatalf("messages count = %d, want 4; output=%s", got, out) + } + if messages[0].Get("tool_calls").Exists() { + t.Fatalf("messages.0 unexpectedly contains tool calls; output=%s", out) + } + if got := messages[1].Get("role").String(); got != "user" { + t.Fatalf("messages.1 role = %q, want user; output=%s", got, out) + } + if got := messages[2].Get("tool_calls.0.id").String(); got != "call_next" { + t.Fatalf("messages.2 tool call ID = %q, want call_next; output=%s", got, out) + } +} + +func TestConvertOpenAIResponsesRequestToOpenAIChatCompletions_MergesDistinctReasoningWithinAssistantTurn(t *testing.T) { + raw := []byte(`{ + "input": [ + {"type":"reasoning","summary":[{"type":"summary_text","text":"first"}]}, + {"type":"message","role":"assistant","reasoning_content":"first","content":[{"type":"output_text","text":"working"}]}, + {"type":"reasoning","summary":[{"type":"summary_text","text":"second"}]}, + {"type":"function_call","call_id":"call_reasoning","name":"exec_command","arguments":"{}"} + ] + }`) + + out := ConvertOpenAIResponsesRequestToOpenAIChatCompletions("kimi-k3", raw, false) + + messages := gjson.GetBytes(out, "messages").Array() + if got := len(messages); got != 1 { + t.Fatalf("messages count = %d, want 1; output=%s", got, out) + } + if got := messages[0].Get("reasoning_content").String(); got != "first\n\nsecond" { + t.Fatalf("reasoning_content = %q, want %q; output=%s", got, "first\n\nsecond", out) + } + if got := messages[0].Get("tool_calls.0.id").String(); got != "call_reasoning" { + t.Fatalf("tool call ID = %q, want call_reasoning; output=%s", got, out) + } +} + +func TestConvertOpenAIResponsesRequestToOpenAIChatCompletions_ReplacesUnavailableReasoningWithinAssistantTurn(t *testing.T) { + raw := []byte(`{ + "input": [ + {"type":"reasoning","summary":[]}, + {"type":"message","role":"assistant","reasoning_content":"real reasoning","content":[{"type":"output_text","text":"working"}]}, + {"type":"function_call","call_id":"call_real_reasoning","name":"exec_command","arguments":"{}"} + ] + }`) + + out := ConvertOpenAIResponsesRequestToOpenAIChatCompletions("kimi-k3", raw, false) + + messages := gjson.GetBytes(out, "messages").Array() + if got := len(messages); got != 1 { + t.Fatalf("messages count = %d, want 1; output=%s", got, out) + } + if got := messages[0].Get("reasoning_content").String(); got != "real reasoning" { + t.Fatalf("reasoning_content = %q, want real reasoning; output=%s", got, out) + } +} + func TestConvertOpenAIResponsesRequestToOpenAIChatCompletions_AttachesReasoningToToolCallMessage(t *testing.T) { raw := []byte(`{ "input": [