mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 17:08:33 +08:00
fix(apicompat): narrow malformed tool-call handling
This commit is contained in:
@@ -280,6 +280,7 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
var lastTurnReasoning string
|
||||
mediaByCallID := make(toolOutputMediaByCallID)
|
||||
invalidFunctionCallIDs := make(map[string]struct{})
|
||||
invalidEmptyFunctionCallOutputs := 0
|
||||
|
||||
reasoningForAssistant := func() string {
|
||||
if pendingReasoning != "" {
|
||||
@@ -342,6 +343,8 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
// turn to self-heal instead of repeatedly replaying the poison.
|
||||
if callID != "" {
|
||||
invalidFunctionCallIDs[callID] = struct{}{}
|
||||
} else {
|
||||
invalidEmptyFunctionCallOutputs++
|
||||
}
|
||||
pendingReasoning = ""
|
||||
continue
|
||||
@@ -403,6 +406,11 @@ func buildChatMessagesFromItems(messages []ChatMessage, rawItems []json.RawMessa
|
||||
case "function_call_output", "custom_tool_call_output", "tool_search_output":
|
||||
outputRaw := bytesTrimSpace(item["output"])
|
||||
callID := rawString(item["call_id"])
|
||||
if callID == "" && invalidEmptyFunctionCallOutputs > 0 {
|
||||
invalidEmptyFunctionCallOutputs--
|
||||
pendingReasoning = ""
|
||||
continue
|
||||
}
|
||||
if _, skipped := invalidFunctionCallIDs[callID]; skipped {
|
||||
pendingReasoning = ""
|
||||
continue
|
||||
@@ -1237,6 +1245,13 @@ func chatMessageToResponsesOutput(message ChatMessage, customTools map[string]bo
|
||||
})
|
||||
continue
|
||||
}
|
||||
// Ordinary Responses function_call arguments must contain valid JSON.
|
||||
// Do not mark a truncated non-streaming Chat tool call as completed;
|
||||
// Codex would persist it and poison the next request in the same way as
|
||||
// the streaming variant guarded by ValidateToolCallArguments.
|
||||
if !json.Valid([]byte(arguments)) {
|
||||
continue
|
||||
}
|
||||
if ns, ok := namespaceTools[toolCall.Function.Name]; ok {
|
||||
outputs = append(outputs, ResponsesOutput{
|
||||
Type: "function_call",
|
||||
@@ -1426,17 +1441,14 @@ func NewChatCompletionsToResponsesStreamState(model string) *ChatCompletionsToRe
|
||||
}
|
||||
|
||||
// ValidateToolCallArguments checks the accumulated function-call arguments
|
||||
// before the stream is finalized. A tool call that ended because the upstream
|
||||
// hit its output limit, or whose SSE chunk was lost, must not be emitted as a
|
||||
// completed Responses item: Codex will persist it and replay it on the next
|
||||
// turn, where a Chat Completions provider rejects the whole request.
|
||||
// before the stream is finalized. A tool call whose argument stream was
|
||||
// truncated must not be emitted as a completed Responses item: Codex will
|
||||
// persist it and replay it on the next turn, where a Chat Completions provider
|
||||
// rejects the whole request.
|
||||
func (state *ChatCompletionsToResponsesStreamState) ValidateToolCallArguments() error {
|
||||
if state == nil {
|
||||
return nil
|
||||
}
|
||||
if state.FinishReason == "length" && len(state.ToolCalls) > 0 {
|
||||
return fmt.Errorf("tool call stream ended at max output length")
|
||||
}
|
||||
for idx, toolCall := range state.ToolCalls {
|
||||
if toolCall == nil {
|
||||
continue
|
||||
|
||||
@@ -37,6 +37,42 @@ func TestResponsesInputToChatMessages_SkipsInvalidHistoricalFunctionCall(t *test
|
||||
require.Equal(t, "user", messages[2].Role)
|
||||
}
|
||||
|
||||
func TestResponsesInputToChatMessages_SkipsInvalidEmptyCallIDOutput(t *testing.T) {
|
||||
input := json.RawMessage(`[
|
||||
{"type":"function_call","call_id":"","name":"exec_command","arguments":"{\"cmd\": \"ssh root@HOST"},
|
||||
{"type":"function_call_output","call_id":"","output":"failed to parse function arguments"},
|
||||
{"role":"user","content":"continue"}
|
||||
]`)
|
||||
|
||||
messages, err := responsesInputToChatMessages("", input)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, messages, 1)
|
||||
require.Equal(t, "user", messages[0].Role)
|
||||
}
|
||||
|
||||
func TestChatCompletionsResponseToResponses_SkipsInvalidFunctionArguments(t *testing.T) {
|
||||
resp := &ChatCompletionsResponse{
|
||||
Model: "deepseek-v4-flash",
|
||||
Choices: []ChatChoice{{
|
||||
Message: ChatMessage{
|
||||
Role: "assistant",
|
||||
ToolCalls: []ChatToolCall{
|
||||
{ID: "call_bad", Type: "function", Function: ChatFunctionCall{Name: "exec_command", Arguments: `{"cmd": "ssh root@HOST`}},
|
||||
{ID: "call_ok", Type: "function", Function: ChatFunctionCall{Name: "exec_command", Arguments: `{}`}},
|
||||
},
|
||||
},
|
||||
FinishReason: "length",
|
||||
}},
|
||||
}
|
||||
|
||||
out := ChatCompletionsResponseToResponses(resp, "deepseek-v4-flash", nil, false, nil)
|
||||
require.Equal(t, "incomplete", out.Status)
|
||||
require.Len(t, out.Output, 1)
|
||||
require.Equal(t, "function_call", out.Output[0].Type)
|
||||
require.Equal(t, "call_ok", out.Output[0].CallID)
|
||||
require.Equal(t, `{}`, out.Output[0].Arguments)
|
||||
}
|
||||
|
||||
func TestResponsesInputToChatMessages_KeepsChatCompletionRoles(t *testing.T) {
|
||||
input := json.RawMessage(`[
|
||||
{"role":"system","content":"system message"},
|
||||
|
||||
@@ -243,7 +243,7 @@ func TestStream_InvalidToolArgumentsAreRejectedBeforeFinalize(t *testing.T) {
|
||||
require.ErrorContains(t, err, "invalid JSON")
|
||||
}
|
||||
|
||||
func TestStream_ToolCallAtOutputLimitIsRejectedBeforeFinalize(t *testing.T) {
|
||||
func TestStream_ValidToolCallAtOutputLimitKeepsIncompleteResponse(t *testing.T) {
|
||||
idx := 0
|
||||
state := NewChatCompletionsToResponsesStreamState("deepseek-v4-flash")
|
||||
chunk := &ChatCompletionsChunk{
|
||||
@@ -269,7 +269,21 @@ func TestStream_ToolCallAtOutputLimitIsRejectedBeforeFinalize(t *testing.T) {
|
||||
ChatCompletionsChunkToResponsesEvents(chunk, state)
|
||||
state.FinishReason = "length"
|
||||
|
||||
require.ErrorContains(t, state.ValidateToolCallArguments(), "max output length")
|
||||
require.NoError(t, state.ValidateToolCallArguments())
|
||||
events := FinalizeChatCompletionsResponsesStream(state)
|
||||
var sawArgsDone, sawIncomplete bool
|
||||
for _, event := range events {
|
||||
switch event.Type {
|
||||
case "response.function_call_arguments.done":
|
||||
sawArgsDone = true
|
||||
require.Equal(t, `{}`, event.Arguments)
|
||||
case "response.completed":
|
||||
require.NotNil(t, event.Response)
|
||||
sawIncomplete = event.Response.Status == "incomplete"
|
||||
}
|
||||
}
|
||||
require.True(t, sawArgsDone)
|
||||
require.True(t, sawIncomplete)
|
||||
}
|
||||
|
||||
// TestStream_SSEWireComplete drives the full stream through SSE encoding and
|
||||
|
||||
@@ -281,13 +281,7 @@ func (s *OpenAIGatewayService) scanCCStream(
|
||||
zap.Error(err),
|
||||
zap.String("request_id", requestID),
|
||||
)
|
||||
// A malformed chunk may contain the middle or the tail of a tool-call
|
||||
// argument. Skipping it and finalizing the accumulated state would turn
|
||||
// a truncated function call into a completed Responses item, which then
|
||||
// poisons Codex's next-turn history. Treat the stream as incomplete so
|
||||
// callers skip finalization and surface an upstream-stream error.
|
||||
st.Err = fmt.Errorf("malformed chat stream chunk: %w", err)
|
||||
break
|
||||
continue
|
||||
}
|
||||
if st.FirstTokenMs == nil && !isOpenAIChatUsageOnlyStreamChunk(payload) && chatChunkStartsResponsesOutput(&chunk) {
|
||||
ms := int(time.Since(startTime).Milliseconds())
|
||||
|
||||
@@ -148,7 +148,7 @@ func TestForwardResponses_ForceChatCompletionsRoutesStreamingToChatCompletions(t
|
||||
require.NotNil(t, result.FirstTokenMs)
|
||||
}
|
||||
|
||||
func TestForwardResponses_ChatFallbackDoesNotFinalizeAfterMalformedStreamChunk(t *testing.T) {
|
||||
func TestForwardResponses_ChatFallbackRejectsInvalidToolArgumentsAtOutputLimit(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-flash","input":"run the command","stream":true}`)
|
||||
@@ -158,42 +158,7 @@ func TestForwardResponses_ChatFallbackDoesNotFinalizeAfterMalformedStreamChunk(t
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
upstreamBody := strings.Join([]string{
|
||||
`data: {"id":"chatcmpl_bad_stream","object":"chat.completion.chunk","model":"deepseek-v4-flash","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_bad","type":"function","function":{"name":"exec_command","arguments":"{}"}}]},"finish_reason":null}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_bad_stream","object":"chat.completion.chunk","choices":[`,
|
||||
"",
|
||||
"data: [DONE]",
|
||||
"",
|
||||
}, "\n")
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "x-request-id": []string{"rid_bad_stream"}},
|
||||
Body: io.NopCloser(strings.NewReader(upstreamBody)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{
|
||||
cfg: rawChatCompletionsTestConfig(),
|
||||
httpUpstream: upstream,
|
||||
}
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, forceChatResponsesFallbackAccount(), body)
|
||||
require.ErrorContains(t, err, "stream usage incomplete")
|
||||
require.NotNil(t, result)
|
||||
require.NotContains(t, rec.Body.String(), "response.function_call_arguments.done")
|
||||
require.NotContains(t, rec.Body.String(), "response.output_item.done")
|
||||
require.NotContains(t, rec.Body.String(), "data: [DONE]")
|
||||
}
|
||||
|
||||
func TestForwardResponses_ChatFallbackDoesNotCompleteToolCallAtOutputLimit(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
body := []byte(`{"model":"deepseek-v4-flash","input":"run the command","stream":true}`)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewReader(body))
|
||||
c.Request.Header.Set("Content-Type", "application/json")
|
||||
|
||||
upstreamBody := strings.Join([]string{
|
||||
`data: {"id":"chatcmpl_length_tool","object":"chat.completion.chunk","model":"deepseek-v4-flash","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_length","type":"function","function":{"name":"exec_command","arguments":"{}"}}]},"finish_reason":null}]}`,
|
||||
`data: {"id":"chatcmpl_length_tool","object":"chat.completion.chunk","model":"deepseek-v4-flash","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_length","type":"function","function":{"name":"exec_command","arguments":"{\"cmd\":\"ssh root@HOST"}}]},"finish_reason":null}]}`,
|
||||
"",
|
||||
`data: {"id":"chatcmpl_length_tool","object":"chat.completion.chunk","model":"deepseek-v4-flash","choices":[{"index":0,"delta":{},"finish_reason":"length"}],"usage":{"prompt_tokens":4,"completion_tokens":6492,"total_tokens":6496}}`,
|
||||
"",
|
||||
@@ -211,8 +176,10 @@ func TestForwardResponses_ChatFallbackDoesNotCompleteToolCallAtOutputLimit(t *te
|
||||
}
|
||||
|
||||
result, err := svc.Forward(context.Background(), c, forceChatResponsesFallbackAccount(), body)
|
||||
require.ErrorContains(t, err, "max output length")
|
||||
require.ErrorContains(t, err, "invalid JSON")
|
||||
require.NotNil(t, result)
|
||||
require.Equal(t, 4, result.Usage.InputTokens)
|
||||
require.Equal(t, 6492, result.Usage.OutputTokens)
|
||||
require.NotContains(t, rec.Body.String(), "response.function_call_arguments.done")
|
||||
require.NotContains(t, rec.Body.String(), "response.output_item.done")
|
||||
require.NotContains(t, rec.Body.String(), "data: [DONE]")
|
||||
|
||||
Reference in New Issue
Block a user