fix(agent): recover direct answer after tool limit

This commit is contained in:
Alix-007 2026-03-17 17:11:57 +08:00
parent dd82794255
commit 4d42697a68
2 changed files with 200 additions and 3 deletions

View file

@ -1505,9 +1505,79 @@ func (al *AgentLoop) runLLMIteration(
}) })
} }
if finalContent == "" && iteration >= agent.MaxIterations && agent.MaxIterations > 0 {
logger.WarnCF("agent", "Tool iteration limit reached, attempting direct answer without tools",
map[string]any{
"agent_id": agent.ID,
"iteration": iteration,
"max_iterations": agent.MaxIterations,
})
llmOpts := map[string]any{
"max_tokens": agent.MaxTokens,
"temperature": agent.Temperature,
"prompt_cache_key": agent.ID,
}
if agent.ThinkingLevel != ThinkingOff {
if tc, ok := agent.Provider.(providers.ThinkingCapable); ok && tc.SupportsThinking() {
llmOpts["thinking_level"] = string(agent.ThinkingLevel)
}
}
directAnswer, recovered, err := al.requestDirectAnswerAfterToolLimit(ctx, agent, messages, activeModel, llmOpts)
switch {
case err != nil:
logger.WarnCF("agent", "Direct-answer fallback after tool limit failed",
map[string]any{
"agent_id": agent.ID,
"iteration": iteration,
"error": err.Error(),
})
case recovered:
finalContent = directAnswer
default:
logger.WarnCF("agent", "Direct-answer fallback produced no usable answer",
map[string]any{
"agent_id": agent.ID,
"iteration": iteration,
})
}
}
return finalContent, iteration, nil return finalContent, iteration, nil
} }
func (al *AgentLoop) requestDirectAnswerAfterToolLimit(
ctx context.Context,
agent *AgentInstance,
messages []providers.Message,
model string,
llmOpts map[string]any,
) (string, bool, error) {
directMessages := append([]providers.Message{}, messages...)
directMessages = append(directMessages, providers.Message{
Role: "user",
Content: "Tool iteration limit reached. Using the available context and tool results already in the conversation, answer the user's latest request directly without calling any more tools.",
})
response, err := agent.Provider.Chat(ctx, directMessages, nil, model, llmOpts)
if err != nil {
return "", false, err
}
if len(response.ToolCalls) > 0 {
return "", false, nil
}
content := response.Content
if content == "" && response.ReasoningContent != "" {
content = response.ReasoningContent
}
if content == "" {
return "", false, nil
}
return content, true, nil
}
// selectCandidates returns the model candidates and resolved model name to use // selectCandidates returns the model candidates and resolved model name to use
// for a conversation turn. When model routing is configured and the incoming // for a conversation turn. When model routing is configured and the incoming
// message scores below the complexity threshold, it returns the light model // message scores below the complexity threshold, it returns the light model

View file

@ -443,6 +443,45 @@ func (m *toolLimitOnlyProvider) GetDefaultModel() string {
return "tool-limit-only-model" return "tool-limit-only-model"
} }
type toolLimitFallbackProvider struct {
calls int
toolsPerCall []int
messagesPerCall [][]providers.Message
}
func (m *toolLimitFallbackProvider) Chat(
ctx context.Context,
messages []providers.Message,
tools []providers.ToolDefinition,
model string,
opts map[string]any,
) (*providers.LLMResponse, error) {
m.calls++
m.toolsPerCall = append(m.toolsPerCall, len(tools))
msgCopy := append([]providers.Message(nil), messages...)
m.messagesPerCall = append(m.messagesPerCall, msgCopy)
if len(tools) > 0 {
return &providers.LLMResponse{
ToolCalls: []providers.ToolCall{{
ID: "call_loop_test",
Type: "function",
Name: "loop_test_tool",
Arguments: map[string]any{"value": "x"},
}},
}, nil
}
return &providers.LLMResponse{
Content: "Recovered direct answer",
ToolCalls: []providers.ToolCall{},
}, nil
}
func (m *toolLimitFallbackProvider) GetDefaultModel() string {
return "tool-limit-fallback-model"
}
// mockCustomTool is a simple mock tool for registration testing // mockCustomTool is a simple mock tool for registration testing
type mockCustomTool struct{} type mockCustomTool struct{}
@ -488,6 +527,29 @@ func (m *toolLimitTestTool) Execute(ctx context.Context, args map[string]any) *t
return tools.SilentResult("tool limit test result") return tools.SilentResult("tool limit test result")
} }
type loopTestTool struct{}
func (m *loopTestTool) Name() string {
return "loop_test_tool"
}
func (m *loopTestTool) Description() string {
return "Loop test tool"
}
func (m *loopTestTool) Parameters() map[string]any {
return map[string]any{
"type": "object",
"properties": map[string]any{
"value": map[string]any{"type": "string"},
},
}
}
func (m *loopTestTool) Execute(ctx context.Context, args map[string]any) *tools.ToolResult {
return tools.SilentResult("loop tool result")
}
// testHelper executes a message and returns the response // testHelper executes a message and returns the response
type testHelper struct { type testHelper struct {
al *AgentLoop al *AgentLoop
@ -1190,6 +1252,57 @@ func TestAgentLoop_ToolLimitUsesDedicatedFallback(t *testing.T) {
if response != toolLimitResponse { if response != toolLimitResponse {
t.Fatalf("response = %q, want %q", response, toolLimitResponse) t.Fatalf("response = %q, want %q", response, toolLimitResponse)
} }
}
func TestAgentLoop_ToolLimitFallsBackToDirectAnswer(t *testing.T) {
tmpDir, err := os.MkdirTemp("", "agent-test-*")
if err != nil {
t.Fatalf("Failed to create temp dir: %v", err)
}
defer os.RemoveAll(tmpDir)
cfg := &config.Config{
Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{
Workspace: tmpDir,
Model: "test-model",
MaxTokens: 4096,
MaxToolIterations: 1,
},
},
}
msgBus := bus.NewMessageBus()
provider := &toolLimitFallbackProvider{}
al := NewAgentLoop(cfg, msgBus, provider)
al.RegisterTool(&loopTestTool{})
sessionKey := "tool-limit-session"
response, err := al.ProcessDirectWithChannel(context.Background(), "hello", sessionKey, "test", "chat1")
if err != nil {
t.Fatalf("ProcessDirectWithChannel failed: %v", err)
}
if response != "Recovered direct answer" {
t.Fatalf("response = %q, want %q", response, "Recovered direct answer")
}
if provider.calls != 2 {
t.Fatalf("provider calls = %d, want 2", provider.calls)
}
if len(provider.toolsPerCall) != 2 {
t.Fatalf("toolsPerCall len = %d, want 2", len(provider.toolsPerCall))
}
if provider.toolsPerCall[0] == 0 {
t.Fatalf("expected first call to include tools, got %v", provider.toolsPerCall)
}
if provider.toolsPerCall[1] != 0 {
t.Fatalf("expected second call to disable tools, got %v", provider.toolsPerCall)
}
fallbackMessages := provider.messagesPerCall[1]
lastMessage := fallbackMessages[len(fallbackMessages)-1]
if lastMessage.Role != "user" || !strings.Contains(lastMessage.Content, "Tool iteration limit reached") {
t.Fatalf("unexpected fallback prompt: %+v", lastMessage)
}
defaultAgent := al.registry.GetDefaultAgent() defaultAgent := al.registry.GetDefaultAgent()
if defaultAgent == nil { if defaultAgent == nil {
@ -1207,8 +1320,23 @@ func TestAgentLoop_ToolLimitUsesDedicatedFallback(t *testing.T) {
t.Fatalf("history len = %d, want 4", len(history)) t.Fatalf("history len = %d, want 4", len(history))
} }
assertRoles(t, history, "user", "assistant", "tool", "assistant") assertRoles(t, history, "user", "assistant", "tool", "assistant")
if history[3].Content != toolLimitResponse { if len(history[1].ToolCalls) != 1 || history[1].ToolCalls[0].Name != "loop_test_tool" {
t.Fatalf("final assistant content = %q, want %q", history[3].Content, toolLimitResponse) if len(history[1].ToolCalls) != 1 ||
history[1].ToolCalls[0].Function == nil ||
history[1].ToolCalls[0].Function.Name != "loop_test_tool" {
t.Fatalf("unexpected assistant tool call history: %+v", history[1].ToolCalls)
}
}
if history[2].Content != "loop tool result" {
t.Fatalf("tool result content = %q, want %q", history[2].Content, "loop tool result")
}
if history[3].Content != "Recovered direct answer" {
t.Fatalf("final assistant content = %q, want %q", history[3].Content, "Recovered direct answer")
}
for _, msg := range history {
if strings.Contains(msg.Content, "Tool iteration limit reached") {
t.Fatalf("synthetic fallback prompt leaked into session history: %+v", history)
}
} }
} }
@ -1223,7 +1351,6 @@ func TestProcessDirectWithChannel_TriggersMCPInitialization(t *testing.T) {
} }
defer os.RemoveAll(tmpDir) defer os.RemoveAll(tmpDir)
// Test with MCP enabled but no servers - should not initialize manager
cfg := &config.Config{ cfg := &config.Config{
Agents: config.AgentsConfig{ Agents: config.AgentsConfig{
Defaults: config.AgentDefaults{ Defaults: config.AgentDefaults{