From d49183b6ab5880a2dc8414a013470a49b01f2ab3 Mon Sep 17 00:00:00 2001 From: Yiliu Date: Fri, 27 Feb 2026 02:07:29 +0800 Subject: [PATCH] fix(agent): preserve tool call context in summarization --- pkg/agent/loop.go | 64 +++++++++++++- pkg/agent/loop_test.go | 192 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 254 insertions(+), 2 deletions(-) diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 8fd7328d1..f6fc6616d 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -1117,7 +1117,7 @@ func (al *AgentLoop) summarizeSession(agent *AgentInstance, sessionKey string) { omitted := false for _, m := range toSummarize { - if m.Role != "user" && m.Role != "assistant" { + if m.Role != "user" && m.Role != "assistant" && m.Role != "tool" { continue } msgTokens := len(m.Content) / 2 @@ -1194,7 +1194,10 @@ func (al *AgentLoop) summarizeBatch( } sb.WriteString("\nCONVERSATION:\n") for _, m := range batch { - fmt.Fprintf(&sb, "%s: %s\n", m.Role, m.Content) + for _, line := range formatMessageForSummary(m) { + sb.WriteString(line) + sb.WriteByte('\n') + } } prompt := sb.String() @@ -1215,6 +1218,63 @@ func (al *AgentLoop) summarizeBatch( return response.Content, nil } +func formatMessageForSummary(msg providers.Message) []string { + content := strings.TrimSpace(msg.Content) + + switch msg.Role { + case "assistant": + lines := make([]string, 0, 1+len(msg.ToolCalls)) + if content != "" { + lines = append(lines, fmt.Sprintf("assistant: %s", content)) + } + for _, tc := range msg.ToolCalls { + name := tc.Name + if name == "" && tc.Function != nil { + name = tc.Function.Name + } + if name == "" { + name = "unknown_tool" + } + + args := "{}" + if tc.Function != nil && strings.TrimSpace(tc.Function.Arguments) != "" { + args = tc.Function.Arguments + } else if len(tc.Arguments) > 0 { + if b, err := json.Marshal(tc.Arguments); err == nil { + args = string(b) + } + } + + lines = append(lines, fmt.Sprintf( + "assistant(tool_call id=%s name=%s): %s", + tc.ID, + name, + utils.Truncate(args, 240), + )) + } + + if len(lines) == 0 { + return []string{"assistant:"} + } + return lines + + case "tool": + toolID := msg.ToolCallID + if toolID == "" { + toolID = "unknown" + } + if content == "" { + return []string{fmt.Sprintf("tool(%s):", toolID)} + } + return []string{fmt.Sprintf("tool(%s): %s", toolID, utils.Truncate(content, 320))} + + case "user": + fallthrough + default: + return []string{fmt.Sprintf("%s: %s", msg.Role, content)} + } +} + // estimateTokens estimates the number of tokens in a message list. // Uses a safe heuristic of 2.5 characters per token to account for CJK and other // overheads better than the previous 3 chars/token. diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index 801b6a46e..ad76879c5 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "strings" "testing" "time" @@ -562,6 +563,27 @@ func (m *failFirstMockProvider) GetDefaultModel() string { return "mock-fail-model" } +type captureSummaryProvider struct { + response string + lastMessages []providers.Message +} + +func (m *captureSummaryProvider) Chat( + ctx context.Context, + messages []providers.Message, + tools []providers.ToolDefinition, + model string, + opts map[string]any, +) (*providers.LLMResponse, error) { + m.lastMessages = make([]providers.Message, len(messages)) + copy(m.lastMessages, messages) + return &providers.LLMResponse{Content: m.response}, nil +} + +func (m *captureSummaryProvider) GetDefaultModel() string { + return "capture-summary-model" +} + // TestAgentLoop_ContextExhaustionRetry verify that the agent retries on context errors func TestAgentLoop_ContextExhaustionRetry(t *testing.T) { tmpDir, err := os.MkdirTemp("", "agent-test-*") @@ -850,4 +872,174 @@ func TestHandleReasoning(t *testing.T) { t.Fatal("expected reasoning message to be dropped when bus is full, but it was published") } }) + t.Run("fallback to default target channel id when disabled", func(t *testing.T) { + al, msgBus := newLoop(t) + cfg := &config.Config{ + Channels: config.ChannelsConfig{ + Telegram: config.TelegramConfig{ReasoningChannelID: "rid-telegram"}, + }, + } + + chManager, err := channels.NewManager(cfg, bus.NewMessageBus(), nil) + if err != nil { + t.Fatalf("Failed to create channel manager: %v", err) + } + chManager.RegisterChannel("telegram", &fakeChannel{id: "rid-telegram"}) + al.cfg = cfg + al.SetChannelManager(chManager) + + al.handleReasoning(context.Background(), "reasoning fallback", "telegram", "rid-telegram") + + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + msg, ok := msgBus.SubscribeOutbound(ctx) + if !ok { + t.Fatal("expected outbound message") + } + if msg.Channel != "telegram" { + t.Fatalf("expected telegram channel, got %+v", msg) + } + if msg.ChatID != "rid-telegram" { + t.Fatalf("expected fallback chat id rid-telegram, got %+v", msg) + } + if msg.Content != "reasoning fallback" { + t.Fatalf("content mismatch: got %q", msg.Content) + } + }) +} + +func TestSummarizeBatch_IncludesToolCallsAndToolResults(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "agent-summary-test-*") + if err != nil { + t.Fatalf("Failed to create temp dir: %v", err) + } + defer os.RemoveAll(tmpDir) + + provider := &captureSummaryProvider{response: "summary ok"} + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{Workspace: tmpDir, Model: "test-model", MaxTokens: 4096}, + }, + } + al := NewAgentLoop(cfg, bus.NewMessageBus(), provider) + agent := al.registry.GetDefaultAgent() + if agent == nil { + t.Fatal("No default agent found") + } + + batch := []providers.Message{ + {Role: "user", Content: "Hi, where are we?"}, + { + Role: "assistant", + ToolCalls: []providers.ToolCall{{ + ID: "call_1", + Function: &providers.FunctionCall{Name: "list_dir", Arguments: `{"path":"."}`}, + }}, + }, + {Role: "tool", ToolCallID: "call_1", Content: "[\"AGENTS.md\",\"README.md\"]"}, + {Role: "assistant", Content: "You're in the workspace root."}, + } + + _, err = al.summarizeBatch(context.Background(), agent, batch, "") + if err != nil { + t.Fatalf("summarizeBatch failed: %v", err) + } + + if len(provider.lastMessages) != 1 { + t.Fatalf("Expected exactly one summary prompt message, got %d", len(provider.lastMessages)) + } + + prompt := provider.lastMessages[0].Content + if !strings.Contains(prompt, "assistant(tool_call id=call_1 name=list_dir):") { + t.Fatalf("Expected tool call serialization in prompt, got: %s", prompt) + } + if !strings.Contains(prompt, "tool(call_1): [\"AGENTS.md\",\"README.md\"]") { + t.Fatalf("Expected tool result serialization in prompt, got: %s", prompt) + } +} + +func TestSummarizeSession_KeepsToolMessagesInSummaryInput(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "agent-summary-session-test-*") + if err != nil { + t.Fatalf("Failed to create temp dir: %v", err) + } + defer os.RemoveAll(tmpDir) + + provider := &captureSummaryProvider{response: "session summary ok"} + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{Workspace: tmpDir, Model: "test-model", MaxTokens: 4096}, + }, + } + al := NewAgentLoop(cfg, bus.NewMessageBus(), provider) + agent := al.registry.GetDefaultAgent() + if agent == nil { + t.Fatal("No default agent found") + } + + sessionKey := "summary-session" + agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "user", Content: "Hi"}) + agent.Sessions.AddFullMessage(sessionKey, providers.Message{ + Role: "assistant", + ToolCalls: []providers.ToolCall{{ + ID: "call_1", + Function: &providers.FunctionCall{Name: "list_dir", Arguments: `{"path":"."}`}, + }}, + }) + agent.Sessions.AddFullMessage( + sessionKey, + providers.Message{Role: "tool", ToolCallID: "call_1", Content: "[\"a\",\"b\"]"}, + ) + agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "assistant", Content: "Done."}) + agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "user", Content: "tail-1"}) + agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "assistant", Content: "tail-2"}) + agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "user", Content: "tail-3"}) + agent.Sessions.AddFullMessage(sessionKey, providers.Message{Role: "assistant", Content: "tail-4"}) + + al.summarizeSession(agent, sessionKey) + + if len(provider.lastMessages) == 0 { + t.Fatal("Expected summarizeSession to call provider") + } + prompt := provider.lastMessages[0].Content + if !strings.Contains(prompt, "tool(call_1): [\"a\",\"b\"]") { + t.Fatalf("Expected tool message preserved in summary prompt, got: %s", prompt) + } +} + +func TestSummarizeBatch_MarksTruncatedToolOutput(t *testing.T) { + tmpDir, err := os.MkdirTemp("", "agent-summary-truncation-test-*") + if err != nil { + t.Fatalf("Failed to create temp dir: %v", err) + } + defer os.RemoveAll(tmpDir) + + provider := &captureSummaryProvider{response: "summary ok"} + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{Workspace: tmpDir, Model: "test-model", MaxTokens: 4096}, + }, + } + al := NewAgentLoop(cfg, bus.NewMessageBus(), provider) + agent := al.registry.GetDefaultAgent() + if agent == nil { + t.Fatal("No default agent found") + } + + longToolOutput := strings.Repeat("file.txt\n", 120) + batch := []providers.Message{{Role: "tool", ToolCallID: "call_1", Content: longToolOutput}} + + _, err = al.summarizeBatch(context.Background(), agent, batch, "") + if err != nil { + t.Fatalf("summarizeBatch failed: %v", err) + } + + if len(provider.lastMessages) != 1 { + t.Fatalf("Expected exactly one summary prompt message, got %d", len(provider.lastMessages)) + } + + prompt := provider.lastMessages[0].Content + if !strings.Contains(prompt, "[TRUNCATED]") { + t.Fatalf("Expected truncation marker in prompt, got: %s", prompt) + } }