diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index ef2951365..6aa99a308 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -461,7 +461,7 @@ func (al *AgentLoop) Run(ctx context.Context) error { if target == nil { cancelDrain() if finalResponse != "" { - al.PublishResponseIfNeeded(ctx, msg.Channel, msg.ChatID, finalResponse) + al.publishResponseIfNeeded(ctx, msg.SessionKey, msg.Channel, msg.ChatID, finalResponse) } return } @@ -521,7 +521,7 @@ func (al *AgentLoop) Run(ctx context.Context) error { } if finalResponse != "" { - al.PublishResponseIfNeeded(ctx, target.Channel, target.ChatID, finalResponse) + al.publishResponseIfNeeded(ctx, target.SessionKey, target.Channel, target.ChatID, finalResponse) } }() } @@ -604,14 +604,17 @@ func (al *AgentLoop) Stop() { } func (al *AgentLoop) PublishResponseIfNeeded(ctx context.Context, channel, chatID, response string) { + al.publishResponseIfNeeded(ctx, "", channel, chatID, response) +} + +func (al *AgentLoop) publishResponseIfNeeded(ctx context.Context, sessionKey, channel, chatID, response string) { if response == "" { return } alreadySent := false - defaultAgent := al.GetRegistry().GetDefaultAgent() - if defaultAgent != nil { - if tool, ok := defaultAgent.Tools.Get("message"); ok { + if agent := al.agentForSession(sessionKey); agent != nil { + if tool, ok := agent.Tools.Get("message"); ok { if mt, ok := tool.(*tools.MessageTool); ok { alreadySent = mt.HasSentInRound() } @@ -1501,11 +1504,7 @@ func (al *AgentLoop) runAgentLoop( } if opts.SendResponse && result.finalContent != "" { - al.bus.PublishOutbound(ctx, bus.OutboundMessage{ - Channel: opts.Channel, - ChatID: opts.ChatID, - Content: result.finalContent, - }) + al.publishResponseIfNeeded(ctx, opts.SessionKey, opts.Channel, opts.ChatID, result.finalContent) } if result.finalContent != "" { diff --git a/pkg/agent/loop_test.go b/pkg/agent/loop_test.go index 25d20c689..e07bc16cb 100644 --- a/pkg/agent/loop_test.go +++ b/pkg/agent/loop_test.go @@ -782,6 +782,83 @@ func TestProcessMessage_MediaArtifactCanBeForwardedBySendFile(t *testing.T) { } } +func TestPublishResponseIfNeeded_SkipsDuplicateForNamedAgentMessageTool(t *testing.T) { + tmpDir := t.TempDir() + cfg := config.DefaultConfig() + cfg.Agents.Defaults.Workspace = tmpDir + cfg.Agents.Defaults.ModelName = "test-model" + cfg.Agents.Defaults.MaxTokens = 4096 + cfg.Agents.Defaults.MaxToolIterations = 10 + cfg.Agents.Defaults.ToolFeedback.Enabled = false + cfg.Agents.List = []config.AgentConfig{ + {ID: "main", Default: true}, + {ID: "support"}, + } + cfg.Bindings = []config.AgentBinding{ + { + AgentID: "support", + Match: config.BindingMatch{ + Channel: "support", + }, + }, + } + + msgBus := bus.NewMessageBus() + provider := &messageToolThenReplyProvider{} + al := NewAgentLoop(cfg, msgBus, provider) + + msg := bus.InboundMessage{ + Channel: "support", + ChatID: "chat1", + SenderID: "user1", + Content: "help", + Peer: bus.Peer{ + Kind: "direct", + ID: "user1", + }, + } + + route, agent, err := al.resolveMessageRoute(msg) + if err != nil { + t.Fatalf("resolveMessageRoute() error = %v", err) + } + if agent.ID != "support" { + t.Fatalf("agent.ID = %q, want %q", agent.ID, "support") + } + + response, err := al.processMessage(context.Background(), msg) + if err != nil { + t.Fatalf("processMessage() error = %v", err) + } + if response != "Named agent final response" { + t.Fatalf("response = %q, want %q", response, "Named agent final response") + } + if provider.calls != 2 { + t.Fatalf("provider calls = %d, want 2", provider.calls) + } + + select { + case outbound := <-msgBus.OutboundChan(): + if outbound.Channel != msg.Channel || outbound.ChatID != msg.ChatID { + t.Fatalf("unexpected outbound target: %+v", outbound) + } + if outbound.Content != "Tool-delivered reply" { + t.Fatalf("outbound content = %q, want %q", outbound.Content, "Tool-delivered reply") + } + case <-time.After(200 * time.Millisecond): + t.Fatal("expected message tool to publish outbound reply") + } + + sessionKey := resolveScopeKey(route, msg.SessionKey) + al.publishResponseIfNeeded(context.Background(), sessionKey, msg.Channel, msg.ChatID, response) + + select { + case extra := <-msgBus.OutboundChan(): + t.Fatalf("expected final publish to be skipped after named-agent message tool send, got %+v", extra) + case <-time.After(50 * time.Millisecond): + } +} + // TestAgentLoop_GetStartupInfo verifies startup info contains tools func TestAgentLoop_GetStartupInfo(t *testing.T) { tmpDir, err := os.MkdirTemp("", "agent-test-*") @@ -1053,6 +1130,40 @@ func (m *toolFeedbackProvider) GetDefaultModel() string { return "heartbeat-tool-feedback-model" } +type messageToolThenReplyProvider struct { + calls int +} + +func (m *messageToolThenReplyProvider) Chat( + ctx context.Context, + messages []providers.Message, + tools []providers.ToolDefinition, + model string, + opts map[string]any, +) (*providers.LLMResponse, error) { + m.calls++ + if m.calls == 1 { + return &providers.LLMResponse{ + Content: "Sending direct response.", + ToolCalls: []providers.ToolCall{{ + ID: "call_named_agent_message", + Type: "function", + Name: "message", + Arguments: map[string]any{"content": "Tool-delivered reply"}, + }}, + }, nil + } + + return &providers.LLMResponse{ + Content: "Named agent final response", + ToolCalls: []providers.ToolCall{}, + }, nil +} + +func (m *messageToolThenReplyProvider) GetDefaultModel() string { + return "message-tool-then-reply-model" +} + type toolLimitOnlyProvider struct{} func (m *toolLimitOnlyProvider) Chat( @@ -1076,6 +1187,39 @@ func (m *toolLimitOnlyProvider) GetDefaultModel() string { return "tool-limit-only-model" } +type messageThenFinalProvider struct { + calls int +} + +func (m *messageThenFinalProvider) Chat( + ctx context.Context, + messages []providers.Message, + tools []providers.ToolDefinition, + model string, + opts map[string]any, +) (*providers.LLMResponse, error) { + m.calls++ + if m.calls == 1 { + return &providers.LLMResponse{ + ToolCalls: []providers.ToolCall{{ + ID: "call_message", + Type: "function", + Name: "message", + Arguments: map[string]any{"content": "sent via message tool"}, + }}, + }, nil + } + + return &providers.LLMResponse{ + Content: "final answer that should be suppressed", + ToolCalls: []providers.ToolCall{}, + }, nil +} + +func (m *messageThenFinalProvider) GetDefaultModel() string { + return "message-then-final-model" +} + // mockCustomTool is a simple mock tool for registration testing type mockCustomTool struct{} @@ -2558,6 +2702,95 @@ func TestProcessMessage_PublishesToolFeedbackWhenEnabled(t *testing.T) { } } +func TestRun_NamedAgentMessageToolSuppressesFinalOutbound(t *testing.T) { + tmpDir := t.TempDir() + cfg := &config.Config{ + Agents: config.AgentsConfig{ + Defaults: config.AgentDefaults{ + Workspace: tmpDir, + ModelName: "test-model", + MaxTokens: 4096, + MaxToolIterations: 10, + }, + List: []config.AgentConfig{ + {ID: "main", Default: true}, + {ID: "worker"}, + }, + }, + Bindings: []config.AgentBinding{ + { + AgentID: "worker", + Match: config.BindingMatch{ + Channel: "test", + AccountID: "*", + Peer: &config.PeerMatch{Kind: "direct", ID: "user-named"}, + }, + }, + }, + Tools: config.ToolsConfig{ + Message: config.ToolConfig{Enabled: true}, + }, + } + + msgBus := bus.NewMessageBus() + provider := &messageThenFinalProvider{} + al := NewAgentLoop(cfg, msgBus, provider) + + runCtx, cancelRun := context.WithCancel(context.Background()) + defer cancelRun() + + runErrCh := make(chan error, 1) + go func() { + runErrCh <- al.Run(runCtx) + }() + + msg := bus.InboundMessage{ + Channel: "test", + SenderID: "user-named", + ChatID: "chat-1", + Content: "trigger named agent tool send", + Peer: bus.Peer{ + Kind: "direct", + ID: "user-named", + }, + } + + pubCtx, pubCancel := context.WithTimeout(context.Background(), 2*time.Second) + defer pubCancel() + if err := msgBus.PublishInbound(pubCtx, msg); err != nil { + t.Fatalf("PublishInbound() error = %v", err) + } + + select { + case outbound := <-msgBus.OutboundChan(): + if outbound.Content != "sent via message tool" { + t.Fatalf("first outbound content = %q, want message tool content", outbound.Content) + } + case <-time.After(2 * time.Second): + t.Fatal("expected message tool outbound") + } + + select { + case outbound := <-msgBus.OutboundChan(): + t.Fatalf("expected final outbound to be suppressed, got %q", outbound.Content) + case <-time.After(300 * time.Millisecond): + } + + if provider.calls < 2 { + t.Fatalf("provider calls = %d, want at least 2 (tool call + final response)", provider.calls) + } + + cancelRun() + select { + case err := <-runErrCh: + if err != nil { + t.Fatalf("Run() returned error: %v", err) + } + case <-time.After(2 * time.Second): + t.Fatal("timeout waiting for Run() to stop") + } +} + func TestResolveMediaRefs_ResolvesToBase64(t *testing.T) { store := media.NewFileMediaStore() dir := t.TempDir()