feat: add cross-channel awareness for AI agent messaging
Enable the AI to send messages across channels (e.g. WS→Discord, Discord→App) by teaching it about available channels and resolving target chat IDs automatically. - Add StateResolver interface and per-channel chatID tracking in state - Make MessageTool Parameters() dynamic with enabled channel list - Resolve "app" alias to last known Android WS session - Resolve cross-channel chatID from state instead of leaking sender's - Only set sentInRound for same-channel sends (fix response suppression) - Add Connected Channels section to system prompt - Wire everything through SetChannelManager and NewAgentLoop Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
7194acbf02
commit
5069da9641
4 changed files with 175 additions and 12 deletions
|
|
@ -16,12 +16,13 @@ import (
|
||||||
)
|
)
|
||||||
|
|
||||||
type ContextBuilder struct {
|
type ContextBuilder struct {
|
||||||
workspace string
|
workspace string
|
||||||
dataDir string
|
dataDir string
|
||||||
skillsLoader *skills.SkillsLoader
|
skillsLoader *skills.SkillsLoader
|
||||||
memory *MemoryStore
|
memory *MemoryStore
|
||||||
tools *tools.ToolRegistry // Direct reference to tool registry
|
tools *tools.ToolRegistry // Direct reference to tool registry
|
||||||
mcpManager *mcp.Manager // MCP server manager
|
mcpManager *mcp.Manager // MCP server manager
|
||||||
|
enabledChannels []string // Active communication channels
|
||||||
}
|
}
|
||||||
|
|
||||||
func getGlobalConfigDir() string {
|
func getGlobalConfigDir() string {
|
||||||
|
|
@ -67,6 +68,11 @@ func (cb *ContextBuilder) SetMCPManager(manager *mcp.Manager) {
|
||||||
cb.mcpManager = manager
|
cb.mcpManager = manager
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetEnabledChannels sets the list of active communication channels for system prompt.
|
||||||
|
func (cb *ContextBuilder) SetEnabledChannels(channels []string) {
|
||||||
|
cb.enabledChannels = channels
|
||||||
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) getIdentity() string {
|
func (cb *ContextBuilder) getIdentity() string {
|
||||||
now := time.Now().Format("2006-01-02 15:04 (Monday)")
|
now := time.Now().Format("2006-01-02 15:04 (Monday)")
|
||||||
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
workspacePath, _ := filepath.Abs(filepath.Join(cb.workspace))
|
||||||
|
|
@ -104,6 +110,21 @@ Your workspace is at: %s
|
||||||
now, runtime, workspacePath, toolsSection)
|
now, runtime, workspacePath, toolsSection)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (cb *ContextBuilder) buildChannelsSection() string {
|
||||||
|
if len(cb.enabledChannels) == 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
var sb strings.Builder
|
||||||
|
sb.WriteString("## Connected Channels\n\n")
|
||||||
|
sb.WriteString("You can send messages to any of these channels using the message tool:\n")
|
||||||
|
for _, ch := range cb.enabledChannels {
|
||||||
|
sb.WriteString(fmt.Sprintf("- %s\n", ch))
|
||||||
|
}
|
||||||
|
sb.WriteString("- app (alias for the current Android app WebSocket session)\n")
|
||||||
|
return sb.String()
|
||||||
|
}
|
||||||
|
|
||||||
func (cb *ContextBuilder) buildToolsSection() string {
|
func (cb *ContextBuilder) buildToolsSection() string {
|
||||||
if cb.tools == nil {
|
if cb.tools == nil {
|
||||||
return ""
|
return ""
|
||||||
|
|
@ -161,6 +182,12 @@ Use the mcp tool to discover and call server tools.
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Connected channels
|
||||||
|
channelsSection := cb.buildChannelsSection()
|
||||||
|
if channelsSection != "" {
|
||||||
|
parts = append(parts, channelsSection)
|
||||||
|
}
|
||||||
|
|
||||||
// Memory context
|
// Memory context
|
||||||
memoryContext := cb.memory.GetMemoryContext()
|
memoryContext := cb.memory.GetMemoryContext()
|
||||||
if memoryContext != "" {
|
if memoryContext != "" {
|
||||||
|
|
|
||||||
|
|
@ -118,6 +118,7 @@ func createToolRegistry(workspace string, restrict bool, cfg *config.Config, msg
|
||||||
})
|
})
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
// StateResolver is injected later in NewAgentLoop after stateManager is created
|
||||||
registry.Register(messageTool)
|
registry.Register(messageTool)
|
||||||
|
|
||||||
return registry
|
return registry
|
||||||
|
|
@ -154,6 +155,18 @@ func NewAgentLoop(cfg *config.Config, msgBus *bus.MessageBus, provider providers
|
||||||
// Create state manager for atomic state persistence
|
// Create state manager for atomic state persistence
|
||||||
stateManager := state.NewManager(dataDir)
|
stateManager := state.NewManager(dataDir)
|
||||||
|
|
||||||
|
// Inject state resolver into message tools for cross-channel "app" alias
|
||||||
|
if tool, ok := toolsRegistry.Get("message"); ok {
|
||||||
|
if mt, ok := tool.(*tools.MessageTool); ok {
|
||||||
|
mt.SetStateResolver(stateManager)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if tool, ok := subagentTools.Get("message"); ok {
|
||||||
|
if mt, ok := tool.(*tools.MessageTool); ok {
|
||||||
|
mt.SetStateResolver(stateManager)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Create context builder and set tools registry
|
// Create context builder and set tools registry
|
||||||
contextBuilder := NewContextBuilder(workspace, dataDir)
|
contextBuilder := NewContextBuilder(workspace, dataDir)
|
||||||
contextBuilder.SetToolsRegistry(toolsRegistry)
|
contextBuilder.SetToolsRegistry(toolsRegistry)
|
||||||
|
|
@ -295,6 +308,17 @@ func (al *AgentLoop) RegisterTool(tool tools.Tool) {
|
||||||
|
|
||||||
func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
|
func (al *AgentLoop) SetChannelManager(cm *channels.Manager) {
|
||||||
al.channelManager = cm
|
al.channelManager = cm
|
||||||
|
|
||||||
|
// Propagate enabled channels to context builder and message tools
|
||||||
|
if cm != nil {
|
||||||
|
channels := cm.GetEnabledChannels()
|
||||||
|
al.contextBuilder.SetEnabledChannels(channels)
|
||||||
|
if tool, ok := al.tools.Get("message"); ok {
|
||||||
|
if mt, ok := tool.(*tools.MessageTool); ok {
|
||||||
|
mt.SetEnabledChannels(channels)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// StateManager returns the state manager used by this agent loop.
|
// StateManager returns the state manager used by this agent loop.
|
||||||
|
|
@ -478,6 +502,10 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (str
|
||||||
logger.WarnCF("agent", "Failed to record last channel: %v", map[string]interface{}{"error": err.Error()})
|
logger.WarnCF("agent", "Failed to record last channel: %v", map[string]interface{}{"error": err.Error()})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
// Record per-channel chatID for cross-channel messaging
|
||||||
|
if err := al.state.SetChannelChatID(opts.Channel, opts.ChatID); err != nil {
|
||||||
|
logger.WarnCF("agent", "Failed to record channel chatID", map[string]interface{}{"error": err.Error()})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"log"
|
"log"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
@ -23,6 +24,10 @@ type State struct {
|
||||||
// Used by heartbeat to always target the main (Android app) session.
|
// Used by heartbeat to always target the main (Android app) session.
|
||||||
LastMainChannel string `json:"last_main_channel,omitempty"`
|
LastMainChannel string `json:"last_main_channel,omitempty"`
|
||||||
|
|
||||||
|
// ChannelChatIDs maps each channel name to the last known chatID.
|
||||||
|
// Used for cross-channel messaging (e.g. WS user sending to Discord).
|
||||||
|
ChannelChatIDs map[string]string `json:"channel_chat_ids,omitempty"`
|
||||||
|
|
||||||
// Timestamp is the last time this state was updated
|
// Timestamp is the last time this state was updated
|
||||||
Timestamp time.Time `json:"timestamp"`
|
Timestamp time.Time `json:"timestamp"`
|
||||||
}
|
}
|
||||||
|
|
@ -79,6 +84,9 @@ func (sm *Manager) SetLastChannel(channel string) error {
|
||||||
sm.state.LastChannel = channel
|
sm.state.LastChannel = channel
|
||||||
sm.state.Timestamp = time.Now()
|
sm.state.Timestamp = time.Now()
|
||||||
|
|
||||||
|
// Also update per-channel chatID mapping (channel format: "name:chatID")
|
||||||
|
sm.updateChannelChatID(channel)
|
||||||
|
|
||||||
// Atomic save using temp file + rename
|
// Atomic save using temp file + rename
|
||||||
if err := sm.saveAtomic(); err != nil {
|
if err := sm.saveAtomic(); err != nil {
|
||||||
return fmt.Errorf("failed to save state atomically: %w", err)
|
return fmt.Errorf("failed to save state atomically: %w", err)
|
||||||
|
|
@ -132,6 +140,9 @@ func (sm *Manager) SetLastChannelWithType(channel, clientType string) error {
|
||||||
sm.state.LastMainChannel = channel
|
sm.state.LastMainChannel = channel
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Also update per-channel chatID mapping (channel format: "name:chatID")
|
||||||
|
sm.updateChannelChatID(channel)
|
||||||
|
|
||||||
if err := sm.saveAtomic(); err != nil {
|
if err := sm.saveAtomic(); err != nil {
|
||||||
return fmt.Errorf("failed to save state atomically: %w", err)
|
return fmt.Errorf("failed to save state atomically: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -146,6 +157,46 @@ func (sm *Manager) GetLastMainChannel() string {
|
||||||
return sm.state.LastMainChannel
|
return sm.state.LastMainChannel
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// updateChannelChatID parses "name:chatID" and updates ChannelChatIDs.
|
||||||
|
// Must be called with the lock held.
|
||||||
|
func (sm *Manager) updateChannelChatID(channelKey string) {
|
||||||
|
if parts := strings.SplitN(channelKey, ":", 2); len(parts) == 2 {
|
||||||
|
if sm.state.ChannelChatIDs == nil {
|
||||||
|
sm.state.ChannelChatIDs = make(map[string]string)
|
||||||
|
}
|
||||||
|
sm.state.ChannelChatIDs[parts[0]] = parts[1]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetChannelChatID records the last known chatID for a given channel name.
|
||||||
|
// This enables cross-channel messaging by resolving the target chatID
|
||||||
|
// when the AI specifies a channel but no chatID.
|
||||||
|
func (sm *Manager) SetChannelChatID(channel, chatID string) error {
|
||||||
|
sm.mu.Lock()
|
||||||
|
defer sm.mu.Unlock()
|
||||||
|
|
||||||
|
if sm.state.ChannelChatIDs == nil {
|
||||||
|
sm.state.ChannelChatIDs = make(map[string]string)
|
||||||
|
}
|
||||||
|
sm.state.ChannelChatIDs[channel] = chatID
|
||||||
|
sm.state.Timestamp = time.Now()
|
||||||
|
|
||||||
|
if err := sm.saveAtomic(); err != nil {
|
||||||
|
return fmt.Errorf("failed to save state atomically: %w", err)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetChannelChatID returns the last known chatID for the given channel name.
|
||||||
|
func (sm *Manager) GetChannelChatID(channel string) string {
|
||||||
|
sm.mu.RLock()
|
||||||
|
defer sm.mu.RUnlock()
|
||||||
|
if sm.state.ChannelChatIDs == nil {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return sm.state.ChannelChatIDs[channel]
|
||||||
|
}
|
||||||
|
|
||||||
// GetTimestamp returns the timestamp of the last state update.
|
// GetTimestamp returns the timestamp of the last state update.
|
||||||
func (sm *Manager) GetTimestamp() time.Time {
|
func (sm *Manager) GetTimestamp() time.Time {
|
||||||
sm.mu.RLock()
|
sm.mu.RLock()
|
||||||
|
|
|
||||||
|
|
@ -3,15 +3,24 @@ package tools
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"strings"
|
||||||
)
|
)
|
||||||
|
|
||||||
type SendCallback func(channel, chatID, content string) error
|
type SendCallback func(channel, chatID, content string) error
|
||||||
|
|
||||||
|
// StateResolver provides access to persistent state for cross-channel routing.
|
||||||
|
type StateResolver interface {
|
||||||
|
GetLastMainChannel() string
|
||||||
|
GetChannelChatID(channel string) string
|
||||||
|
}
|
||||||
|
|
||||||
type MessageTool struct {
|
type MessageTool struct {
|
||||||
sendCallback SendCallback
|
sendCallback SendCallback
|
||||||
defaultChannel string
|
defaultChannel string
|
||||||
defaultChatID string
|
defaultChatID string
|
||||||
sentInRound bool // Tracks whether a message was sent in the current processing round
|
sentInRound bool // Tracks whether a message was sent in the current processing round
|
||||||
|
enabledChannels []string
|
||||||
|
stateResolver StateResolver
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMessageTool() *MessageTool {
|
func NewMessageTool() *MessageTool {
|
||||||
|
|
@ -27,6 +36,11 @@ func (t *MessageTool) Description() string {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *MessageTool) Parameters() map[string]interface{} {
|
func (t *MessageTool) Parameters() map[string]interface{} {
|
||||||
|
channelDesc := "Optional: target channel (telegram, whatsapp, etc.)"
|
||||||
|
if len(t.enabledChannels) > 0 {
|
||||||
|
channelDesc = fmt.Sprintf("Target channel. Available: %s, app (= current Android app session). Omit to reply on the current channel.",
|
||||||
|
strings.Join(t.enabledChannels, ", "))
|
||||||
|
}
|
||||||
return map[string]interface{}{
|
return map[string]interface{}{
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": map[string]interface{}{
|
"properties": map[string]interface{}{
|
||||||
|
|
@ -36,7 +50,7 @@ func (t *MessageTool) Parameters() map[string]interface{} {
|
||||||
},
|
},
|
||||||
"channel": map[string]interface{}{
|
"channel": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "Optional: target channel (telegram, whatsapp, etc.)",
|
"description": channelDesc,
|
||||||
},
|
},
|
||||||
"chat_id": map[string]interface{}{
|
"chat_id": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
|
|
@ -47,6 +61,16 @@ func (t *MessageTool) Parameters() map[string]interface{} {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetEnabledChannels updates the list of available channel names for parameter descriptions.
|
||||||
|
func (t *MessageTool) SetEnabledChannels(channels []string) {
|
||||||
|
t.enabledChannels = channels
|
||||||
|
}
|
||||||
|
|
||||||
|
// SetStateResolver sets the state resolver for cross-channel alias resolution.
|
||||||
|
func (t *MessageTool) SetStateResolver(sr StateResolver) {
|
||||||
|
t.stateResolver = sr
|
||||||
|
}
|
||||||
|
|
||||||
func (t *MessageTool) SetContext(channel, chatID string) {
|
func (t *MessageTool) SetContext(channel, chatID string) {
|
||||||
t.defaultChannel = channel
|
t.defaultChannel = channel
|
||||||
t.defaultChatID = chatID
|
t.defaultChatID = chatID
|
||||||
|
|
@ -71,9 +95,37 @@ func (t *MessageTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
channel, _ := args["channel"].(string)
|
channel, _ := args["channel"].(string)
|
||||||
chatID, _ := args["chat_id"].(string)
|
chatID, _ := args["chat_id"].(string)
|
||||||
|
|
||||||
|
// Resolve "app" alias to the last known Android app (main) WebSocket session
|
||||||
|
if channel == "app" && t.stateResolver != nil {
|
||||||
|
if mainCh := t.stateResolver.GetLastMainChannel(); mainCh != "" {
|
||||||
|
// mainCh format: "websocket:ws:uuid"
|
||||||
|
parts := strings.SplitN(mainCh, ":", 2)
|
||||||
|
if len(parts) == 2 {
|
||||||
|
channel = parts[0]
|
||||||
|
chatID = parts[1]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if channel == "" {
|
if channel == "" {
|
||||||
channel = t.defaultChannel
|
channel = t.defaultChannel
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Cross-channel send: AI specified a different channel but no chatID.
|
||||||
|
// Look up the last known chatID for the target channel from state,
|
||||||
|
// instead of using the current session's defaultChatID (which belongs
|
||||||
|
// to a different channel and would cause API errors).
|
||||||
|
isCrossChannel := channel != "" && channel != t.defaultChannel
|
||||||
|
if chatID == "" && isCrossChannel && t.stateResolver != nil {
|
||||||
|
chatID = t.stateResolver.GetChannelChatID(channel)
|
||||||
|
}
|
||||||
|
if chatID == "" && isCrossChannel {
|
||||||
|
return &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf("Cannot send to %s: no known chat_id. A message must be received from %s first so the system can learn its chat_id.", channel, channel),
|
||||||
|
IsError: true,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if chatID == "" {
|
if chatID == "" {
|
||||||
chatID = t.defaultChatID
|
chatID = t.defaultChatID
|
||||||
}
|
}
|
||||||
|
|
@ -94,7 +146,12 @@ func (t *MessageTool) Execute(ctx context.Context, args map[string]interface{})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
t.sentInRound = true
|
// Only mark as "sent in round" when the message went to the originating channel.
|
||||||
|
// Cross-channel sends (e.g. WS→Discord) must NOT suppress the response
|
||||||
|
// back to the sender's channel.
|
||||||
|
if channel == t.defaultChannel {
|
||||||
|
t.sentInRound = true
|
||||||
|
}
|
||||||
// Silent: user already received the message directly
|
// Silent: user already received the message directly
|
||||||
return &ToolResult{
|
return &ToolResult{
|
||||||
ForLLM: fmt.Sprintf("Message sent to %s:%s", channel, chatID),
|
ForLLM: fmt.Sprintf("Message sent to %s:%s", channel, chatID),
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue