From ba1376ae6d711a8c929439d4ce3b771a65be92d5 Mon Sep 17 00:00:00 2001 From: Oceanpie Date: Thu, 5 Mar 2026 20:18:47 +0800 Subject: [PATCH] fix(agent): drop assistant tool_calls turns with incomplete tool results --- pkg/agent/context.go | 47 +++++++++++++++++++++++++++++++++++++- pkg/agent/context_test.go | 48 +++++++++++++++++++++++++++++++++++++++ 2 files changed, 94 insertions(+), 1 deletion(-) diff --git a/pkg/agent/context.go b/pkg/agent/context.go index 3aa903b3f..31ab0b8c0 100644 --- a/pkg/agent/context.go +++ b/pkg/agent/context.go @@ -602,7 +602,52 @@ func sanitizeHistoryForProvider(history []providers.Message) []providers.Message } } - return sanitized + final := make([]providers.Message, 0, len(sanitized)) + for i := 0; i < len(sanitized); i++ { + msg := sanitized[i] + if msg.Role != "assistant" || len(msg.ToolCalls) == 0 { + final = append(final, msg) + continue + } + + expected := make(map[string]bool, len(msg.ToolCalls)) + for _, tc := range msg.ToolCalls { + expected[tc.ID] = false + } + + for j := i + 1; j < len(sanitized); j++ { + next := sanitized[j] + if next.Role != "tool" { + break + } + if _, ok := expected[next.ToolCallID]; ok { + expected[next.ToolCallID] = true + } + } + + allFound := true + for _, found := range expected { + if !found { + allFound = false + break + } + } + if !allFound { + logger.DebugCF( + "agent", + "Dropping assistant tool-call turn with incomplete tool results", + map[string]any{"tool_calls": len(expected)}, + ) + for i+1 < len(sanitized) && sanitized[i+1].Role == "tool" { + i++ + } + continue + } + + final = append(final, msg) + } + + return final } func (cb *ContextBuilder) AddToolResult( diff --git a/pkg/agent/context_test.go b/pkg/agent/context_test.go index e023c9c30..036e2867c 100644 --- a/pkg/agent/context_test.go +++ b/pkg/agent/context_test.go @@ -188,6 +188,54 @@ func TestSanitizeHistoryForProvider_PlainConversation(t *testing.T) { assertRoles(t, result, "user", "assistant", "user", "assistant") } +func TestSanitizeHistoryForProvider_AssistantToolCallMissingAllResults(t *testing.T) { + history := []providers.Message{ + msg("user", "run tool"), + assistantWithTools("A"), + msg("assistant", "fallback text"), + } + + result := sanitizeHistoryForProvider(history) + if len(result) != 2 { + t.Fatalf("expected 2 messages, got %d: %+v", len(result), roles(result)) + } + assertRoles(t, result, "user", "assistant") +} + +func TestSanitizeHistoryForProvider_AssistantToolCallIncompleteResults(t *testing.T) { + history := []providers.Message{ + msg("user", "run two tools"), + assistantWithTools("A", "B"), + toolResult("A"), + msg("assistant", "next"), + } + + result := sanitizeHistoryForProvider(history) + if len(result) != 2 { + t.Fatalf("expected 2 messages, got %d: %+v", len(result), roles(result)) + } + assertRoles(t, result, "user", "assistant") +} + +func TestSanitizeHistoryForProvider_DropsBrokenRoundKeepsCompleteRound(t *testing.T) { + history := []providers.Message{ + msg("user", "first"), + assistantWithTools("A", "B"), + toolResult("A"), + msg("assistant", "after broken round"), + msg("user", "second"), + assistantWithTools("C"), + toolResult("C"), + msg("assistant", "done"), + } + + result := sanitizeHistoryForProvider(history) + if len(result) != 6 { + t.Fatalf("expected 6 messages, got %d: %+v", len(result), roles(result)) + } + assertRoles(t, result, "user", "assistant", "user", "assistant", "tool", "assistant") +} + func roles(msgs []providers.Message) []string { r := make([]string, len(msgs)) for i, m := range msgs {