From da6a862079657c34e45af16e6d76508a0a744541 Mon Sep 17 00:00:00 2001 From: Alix-007 <267018309+Alix-007@users.noreply.github.com> Date: Fri, 3 Apr 2026 16:53:29 +0800 Subject: [PATCH] fix(agent): normalize tool-call history for strict providers --- pkg/agent/context.go | 134 +++++++++++++++++++++++--------------- pkg/agent/context_test.go | 65 +++++++++++++----- 2 files changed, 128 insertions(+), 71 deletions(-) diff --git a/pkg/agent/context.go b/pkg/agent/context.go index c2921294b..7c311e79b 100644 --- a/pkg/agent/context.go +++ b/pkg/agent/context.go @@ -681,75 +681,103 @@ func sanitizeHistoryForProvider(history []providers.Message) []providers.Message } } - // Second pass: ensure every assistant message with tool_calls has matching - // tool result messages following it. This is required by strict providers - // like DeepSeek that enforce: "An assistant message with 'tool_calls' must - // be followed by tool messages responding to each 'tool_call_id'." + // Second pass: normalize each assistant tool-call round so strict providers + // always receive exactly one tool response for each expected tool_call_id. + // This avoids cross-round dedupe bugs when providers reuse tool_call IDs. final := make([]providers.Message, 0, len(sanitized)) - seenToolCallID := make(map[string]bool) for i := 0; i < len(sanitized); i++ { msg := sanitized[i] + if msg.Role != "assistant" || len(msg.ToolCalls) == 0 { + final = append(final, msg) + continue + } - // Deduplicate tool results by ToolCallID - if msg.Role == "tool" && msg.ToolCallID != "" { - if seenToolCallID[msg.ToolCallID] { - logger.DebugCF("agent", "Dropping duplicate tool result", map[string]any{ - "tool_call_id": msg.ToolCallID, + // Collect expected IDs in stable order while deduplicating duplicates + // within the same assistant turn. + expectedOrder := make([]string, 0, len(msg.ToolCalls)) + expectedSet := make(map[string]struct{}, len(msg.ToolCalls)) + for _, tc := range msg.ToolCalls { + if tc.ID == "" { + continue + } + if _, exists := expectedSet[tc.ID]; exists { + continue + } + expectedSet[tc.ID] = struct{}{} + expectedOrder = append(expectedOrder, tc.ID) + } + + // If provider output is malformed and tool calls have no IDs, keep the + // assistant turn as-is and let provider-level validation handle it. + if len(expectedOrder) == 0 { + final = append(final, msg) + continue + } + + // Gather the contiguous tool result block for this assistant turn. + j := i + 1 + collectedTools := make([]providers.Message, 0, len(expectedOrder)) + for ; j < len(sanitized); j++ { + if sanitized[j].Role != "tool" { + break + } + collectedTools = append(collectedTools, sanitized[j]) + } + + // Keep only the first tool result per expected ID and drop stray entries. + toolByID := make(map[string]providers.Message, len(expectedOrder)) + for _, toolMsg := range collectedTools { + if toolMsg.ToolCallID == "" { + logger.DebugCF("agent", "Dropping tool result without tool_call_id", map[string]any{}) + continue + } + if _, expected := expectedSet[toolMsg.ToolCallID]; !expected { + logger.DebugCF("agent", "Dropping unexpected tool result", map[string]any{ + "tool_call_id": toolMsg.ToolCallID, }) continue } - seenToolCallID[msg.ToolCallID] = true - } - - if msg.Role == "assistant" && len(msg.ToolCalls) > 0 { - // Collect expected tool_call IDs - expected := make(map[string]bool, len(msg.ToolCalls)) - for _, tc := range msg.ToolCalls { - expected[tc.ID] = false - } - - // Check following messages for matching tool results - toolMsgCount := 0 - for j := i + 1; j < len(sanitized); j++ { - if sanitized[j].Role != "tool" { - break - } - toolMsgCount++ - if _, exists := expected[sanitized[j].ToolCallID]; exists { - expected[sanitized[j].ToolCallID] = true - } - } - - // If any tool_call_id is missing, drop this assistant message and its partial tool messages - allFound := true - for toolCallID, found := range expected { - if !found { - allFound = false - logger.DebugCF( - "agent", - "Dropping assistant message with incomplete tool results", - map[string]any{ - "missing_tool_call_id": toolCallID, - "expected_count": len(expected), - "found_count": toolMsgCount, - }, - ) - break - } - } - - if !allFound { - // Skip this assistant message and its tool messages - i += toolMsgCount + if _, exists := toolByID[toolMsg.ToolCallID]; exists { + logger.DebugCF("agent", "Dropping duplicate tool result within turn", map[string]any{ + "tool_call_id": toolMsg.ToolCallID, + }) continue } + toolByID[toolMsg.ToolCallID] = toolMsg } + final = append(final, msg) + missingCount := 0 + for _, toolCallID := range expectedOrder { + if toolMsg, ok := toolByID[toolCallID]; ok { + final = append(final, toolMsg) + continue + } + missingCount++ + final = append(final, syntheticToolResultForMissingID(toolCallID)) + } + if missingCount > 0 { + logger.DebugCF("agent", "Inserted synthetic tool results for missing IDs", map[string]any{ + "missing_count": missingCount, + "expected_count": len(expectedOrder), + }) + } + + // Skip the collected tool block; it has already been normalized above. + i = j - 1 } return final } +func syntheticToolResultForMissingID(toolCallID string) providers.Message { + return providers.Message{ + Role: "tool", + ToolCallID: toolCallID, + Content: "Tool result unavailable: this call response was missing in retained history.", + } +} + func (cb *ContextBuilder) AddToolResult( messages []providers.Message, toolCallID, toolName, result string, diff --git a/pkg/agent/context_test.go b/pkg/agent/context_test.go index 0d7948eef..94c2c27b7 100644 --- a/pkg/agent/context_test.go +++ b/pkg/agent/context_test.go @@ -249,13 +249,15 @@ func TestSanitizeHistoryForProvider_IncompleteToolResults(t *testing.T) { } result := sanitizeHistoryForProvider(history) - // The assistant message with incomplete tool results should be dropped, - // along with its partial tool result. The remaining messages are: - // user ("do two things"), user ("next question"), assistant ("answer") - if len(result) != 3 { - t.Fatalf("expected 3 messages, got %d: %+v", len(result), roles(result)) + // The incomplete turn should be normalized, not dropped: + // one synthetic tool result is inserted for missing "B". + if len(result) != 6 { + t.Fatalf("expected 6 messages, got %d: %+v", len(result), roles(result)) + } + assertRoles(t, result, "user", "assistant", "tool", "tool", "user", "assistant") + if result[3].ToolCallID != "B" { + t.Fatalf("expected synthetic tool result for B, got %q", result[3].ToolCallID) } - assertRoles(t, result, "user", "user", "assistant") } // TestSanitizeHistoryForProvider_MissingAllToolResults tests the case where @@ -270,12 +272,14 @@ func TestSanitizeHistoryForProvider_MissingAllToolResults(t *testing.T) { } result := sanitizeHistoryForProvider(history) - // The assistant message with no tool results should be dropped. - // Remaining: user ("do something"), user ("hello"), assistant ("hi") - if len(result) != 3 { - t.Fatalf("expected 3 messages, got %d: %+v", len(result), roles(result)) + // No real tool result exists, so one synthetic tool result is inserted. + if len(result) != 5 { + t.Fatalf("expected 5 messages, got %d: %+v", len(result), roles(result)) + } + assertRoles(t, result, "user", "assistant", "tool", "user", "assistant") + if result[2].ToolCallID != "A" { + t.Fatalf("expected synthetic tool result for A, got %q", result[2].ToolCallID) } - assertRoles(t, result, "user", "user", "assistant") } // TestSanitizeHistoryForProvider_PartialToolResultsInMiddle tests that @@ -297,12 +301,37 @@ func TestSanitizeHistoryForProvider_PartialToolResultsInMiddle(t *testing.T) { } result := sanitizeHistoryForProvider(history) - // First round is complete (user, assistant+tools, tool, assistant), - // second round is incomplete and dropped (assistant+tools, partial tool), - // third round is complete (user, assistant+tools, tool, assistant). - // Remaining: user, assistant, tool, assistant, user, user, assistant, tool, assistant - if len(result) != 9 { - t.Fatalf("expected 9 messages, got %d: %+v", len(result), roles(result)) + // First and third rounds remain unchanged; second round gets a synthetic + // tool result for missing "C" instead of being dropped. + if len(result) != 12 { + t.Fatalf("expected 12 messages, got %d: %+v", len(result), roles(result)) + } + assertRoles(t, result, "user", "assistant", "tool", "assistant", "user", "assistant", "tool", "tool", "user", "assistant", "tool", "assistant") + if result[7].ToolCallID != "C" { + t.Fatalf("expected synthetic tool result for C, got %q", result[7].ToolCallID) + } +} + +// TestSanitizeHistoryForProvider_ReusedToolCallIDAcrossRounds ensures +// per-turn normalization does not deduplicate tool_call IDs globally. +func TestSanitizeHistoryForProvider_ReusedToolCallIDAcrossRounds(t *testing.T) { + history := []providers.Message{ + msg("user", "round one"), + assistantWithTools("call_1"), + toolResult("call_1"), + msg("assistant", "done one"), + msg("user", "round two"), + assistantWithTools("call_1"), // ID reused by provider in a later round + toolResult("call_1"), + msg("assistant", "done two"), + } + + result := sanitizeHistoryForProvider(history) + if len(result) != 8 { + t.Fatalf("expected 8 messages, got %d: %+v", len(result), roles(result)) + } + assertRoles(t, result, "user", "assistant", "tool", "assistant", "user", "assistant", "tool", "assistant") + if result[2].ToolCallID != "call_1" || result[6].ToolCallID != "call_1" { + t.Fatalf("expected both rounds to keep tool_call_id call_1, got %q and %q", result[2].ToolCallID, result[6].ToolCallID) } - assertRoles(t, result, "user", "assistant", "tool", "assistant", "user", "user", "assistant", "tool", "assistant") }