fix: detect and break duplicate tool call loops with system nudge (#819)

This commit is contained in:
Rahul Bansal 2026-03-01 22:15:22 +05:30
parent 02a19bfb76
commit 64d1dddb0d
3 changed files with 183 additions and 2 deletions

View file

@ -625,6 +625,10 @@ func (al *AgentLoop) runLLMIteration(
iteration := 0 iteration := 0
var finalContent string var finalContent string
// Duplicate tool call loop detection
var lastToolCallSig string
consecutiveDups := 0
for iteration < agent.MaxIterations { for iteration < agent.MaxIterations {
iteration++ iteration++
@ -808,6 +812,70 @@ func (al *AgentLoop) runLLMIteration(
"iteration": iteration, "iteration": iteration,
}) })
// Duplicate tool call loop detection: build a signature from tool names + arguments
var sigParts []string
for _, tc := range normalizedToolCalls {
argsJSON, _ := json.Marshal(tc.Arguments)
sigParts = append(sigParts, tc.Name+":"+string(argsJSON))
}
currentSig := strings.Join(sigParts, "|")
if currentSig == lastToolCallSig {
consecutiveDups++
} else {
lastToolCallSig = currentSig
consecutiveDups = 1
}
if consecutiveDups >= 3 {
logger.WarnCF("agent", "Duplicate tool call loop detected, injecting nudge",
map[string]any{
"agent_id": agent.ID,
"iteration": iteration,
"consecutive_dups": consecutiveDups,
"signature": currentSig,
})
// Build assistant message with the duplicate tool calls
dupAssistantMsg := providers.Message{
Role: "assistant",
Content: response.Content,
}
for _, tc := range normalizedToolCalls {
argumentsJSON, _ := json.Marshal(tc.Arguments)
dupAssistantMsg.ToolCalls = append(dupAssistantMsg.ToolCalls, providers.ToolCall{
ID: tc.ID,
Type: "function",
Name: tc.Name,
Function: &providers.FunctionCall{
Name: tc.Name,
Arguments: string(argumentsJSON),
},
})
}
messages = append(messages, dupAssistantMsg)
// Add synthetic tool results for each tool call (skipped)
for _, tc := range normalizedToolCalls {
messages = append(messages, providers.Message{
Role: "tool",
Content: "[Skipped: duplicate tool call detected — same call repeated 3+ times]",
ToolCallID: tc.ID,
})
}
// Add a system nudge to break the loop
messages = append(messages, providers.Message{
Role: "user",
Content: "[System] You have been repeating the same tool call. Please try a different approach or provide a final answer.",
})
// Reset counter so the model gets a fresh chance
consecutiveDups = 0
lastToolCallSig = ""
continue
}
// Build assistant message with tool calls // Build assistant message with tool calls
assistantMsg := providers.Message{ assistantMsg := providers.Message{
Role: "assistant", Role: "assistant",
@ -931,6 +999,15 @@ func (al *AgentLoop) runLLMIteration(
} }
} }
if iteration >= agent.MaxIterations {
logger.WarnCF("agent", "Reached max tool iterations",
map[string]any{
"agent_id": agent.ID,
"max": agent.MaxIterations,
"iteration": iteration,
})
}
return finalContent, iteration, nil return finalContent, iteration, nil
} }

View file

@ -787,7 +787,7 @@ func TestHandleReasoning(t *testing.T) {
} }
}) })
t.Run("returns promptly when bus is full", func(t *testing.T) { t.Run("returns promptly when bus is full", func(t *testing.T) { //nolint:dupl
al, msgBus := newLoop(t) al, msgBus := newLoop(t)
// Fill the outbound bus buffer until a publish would block. // Fill the outbound bus buffer until a publish would block.
@ -840,3 +840,102 @@ func TestHandleReasoning(t *testing.T) {
} }
}) })
} }
// dummyTool is a tool that always succeeds, used for duplicate detection testing.
type dummyTool struct {
name string
}
func (d *dummyTool) Name() string { return d.name }
func (d *dummyTool) Description() string { return "dummy tool for testing" }
func (d *dummyTool) Parameters() map[string]any {
return map[string]any{
"type": "object",
"properties": map[string]any{},
}
}
func (d *dummyTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
return tools.SilentResult("ok")
}
func TestRunLLMIteration_DuplicateToolCallBreaker(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-dup-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
callCount := 0
dupToolCall := providers.ToolCall{
ID: "call_dup",
Type: "function",
Name: "dummy_read",
Function: &providers.FunctionCall{
Name: "dummy_read",
Arguments: `{"path":"test.txt"}`,
},
Arguments: map[string]any{"path": "test.txt"},
}
provider := &mockProvider{
chatFunc: func(ctx context.Context, messages []providers.Message, toolDefs []providers.ToolDefinition, model string, opts map[string]any) (*providers.LLMResponse, error) {
callCount++
// First 4 calls return duplicate tool calls, 5th returns text
if callCount <= 4 {
return &providers.LLMResponse{
Content: "",
ToolCalls: []providers.ToolCall{dupToolCall},
}, nil
}
return &providers.LLMResponse{
Content: "Done after breaking loop",
ToolCalls: []providers.ToolCall{},
}, nil
},
}
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "test-model",
MaxTokens: 4096,
MaxToolIterations: 20,
},
},
}
msgBus := bus.NewMessageBus()
al := NewAgentLoop(cfg, msgBus, provider)
agent := al.registry.GetDefaultAgent()
if agent == nil {
t.Fatal("no default agent")
}
agent.Tools.Register(&dummyTool{name: "dummy_read"})
messages := []providers.Message{
{Role: "system", Content: "You are a helpful assistant."},
{Role: "user", Content: "Read test.txt"},
}
resp, _, err := al.runLLMIteration(context.Background(), agent, messages, processOptions{
SessionKey: "test-dup",
Channel: "test",
ChatID: "chat1",
DefaultResponse: "default",
})
if err != nil {
t.Fatalf("runLLMIteration error: %v", err)
}
if resp != "Done after breaking loop" {
t.Errorf("expected 'Done after breaking loop', got %q", resp)
}
// The breaker should trigger after 3 consecutive dupes, meaning
// we should see fewer total LLM calls than MaxIterations
if callCount > 10 {
t.Errorf("expected breaker to limit calls, got %d", callCount)
}
}

View file

@ -6,7 +6,9 @@ import (
"github.com/sipeed/picoclaw/pkg/providers" "github.com/sipeed/picoclaw/pkg/providers"
) )
type mockProvider struct{} type mockProvider struct {
chatFunc func(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, opts map[string]any) (*providers.LLMResponse, error)
}
func (m *mockProvider) Chat( func (m *mockProvider) Chat(
ctx context.Context, ctx context.Context,
@ -15,6 +17,9 @@ func (m *mockProvider) Chat(
model string, model string,
opts map[string]any, opts map[string]any,
) (*providers.LLMResponse, error) { ) (*providers.LLMResponse, error) {
if m.chatFunc != nil {
return m.chatFunc(ctx, messages, tools, model, opts)
}
return &providers.LLMResponse{ return &providers.LLMResponse{
Content: "Mock response", Content: "Mock response",
ToolCalls: []providers.ToolCall{}, ToolCalls: []providers.ToolCall{},