From e6558f3ba747f5a95b5219e22d3e0834d63db6ce Mon Sep 17 00:00:00 2001 From: t <452214834@qq.com> Date: Mon, 9 Mar 2026 15:17:45 +0800 Subject: [PATCH] feat: async command handling for instant /stop response Commands are now processed in separate goroutines, allowing /stop to immediately cancel long-running tasks without waiting in the message queue. Changes: - Add handleCommandAsync method for async command processing - Commands with prefix (/ or !) are handled asynchronously - Non-command messages continue to be processed synchronously - Add mutex-protected currentCancel for thread-safe task cancellation --- pkg/agent/loop.go | 178 +++++++++++++++++++++++++++++++++++++++++----- 1 file changed, 161 insertions(+), 17 deletions(-) diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 72affa020..5840bb2ce 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -48,6 +48,10 @@ type AgentLoop struct { mediaStore media.MediaStore transcriber voice.Transcriber cmdRegistry *commands.Registry + + // Task cancellation support + currentCancel context.CancelFunc + currentCancelMu sync.Mutex } // processOptions configures how a message is processed @@ -85,6 +89,13 @@ func NewAgentLoop( // Initialize conversation logger logger.InitConversationLogger(cfg.Tools.ConversationLog, workspace) + // Initialize sanitizer for sensitive data masking + utils.InitGlobalSanitizer(utils.SanitizerConfig{ + Enabled: cfg.Tools.Sanitizer.Enabled, + Keywords: convertSanitizerKeywords(cfg.Tools.Sanitizer.Keywords), + CustomPatterns: convertSanitizerPatterns(cfg.Tools.Sanitizer.CustomPatterns), + }) + // Wrap provider with logging if enabled llmLogger := logger.GetLLMLogger() if llmLogger != nil && llmLogger.IsEnabled() { @@ -356,24 +367,31 @@ func (al *AgentLoop) Run(ctx context.Context) error { continue } - // Process message - func() { - // TODO: Re-enable media cleanup after inbound media is properly consumed by the agent. - // Currently disabled because files are deleted before the LLM can access their content. - // defer func() { - // if al.mediaStore != nil && msg.MediaScope != "" { - // if releaseErr := al.mediaStore.ReleaseAll(msg.MediaScope); releaseErr != nil { - // logger.WarnCF("agent", "Failed to release media", map[string]any{ - // "scope": msg.MediaScope, - // "error": releaseErr.Error(), - // }) - // } - // } - // }() + // Check if this is a command - commands are handled asynchronously + // so they can interrupt long-running tasks (e.g., /stop) + if commands.HasCommandPrefix(msg.Content) { + go al.handleCommandAsync(ctx, msg) + continue + } - response, err := al.processMessage(ctx, msg) + // Process non-command message synchronously + func() { + // Create cancellable context for this message + msgCtx, msgCancel := context.WithCancel(ctx) + defer msgCancel() + + // Store cancel function for /stop command + al.setCurrentCancel(msgCancel) + defer al.clearCurrentCancel() + + response, err := al.processMessage(msgCtx, msg) if err != nil { - response = fmt.Sprintf("Error processing message: %v", err) + // Check if the error is due to context cancellation (user issued /stop) + if errors.Is(err, context.Canceled) { + response = "⏹️ Task stopped." + } else { + response = fmt.Sprintf("Error processing message: %v", err) + } } if response != "" { @@ -421,6 +439,35 @@ func (al *AgentLoop) Stop() { al.running.Store(false) } +// setCurrentCancel stores the cancel function for the current task. +func (al *AgentLoop) setCurrentCancel(cancel context.CancelFunc) { + al.currentCancelMu.Lock() + defer al.currentCancelMu.Unlock() + al.currentCancel = cancel +} + +// clearCurrentCancel clears the cancel function after task completion. +func (al *AgentLoop) clearCurrentCancel() { + al.currentCancelMu.Lock() + defer al.currentCancelMu.Unlock() + al.currentCancel = nil +} + +// CancelCurrentTask cancels the currently running task, if any. +// Returns true if a task was cancelled, false if no task was running. +func (al *AgentLoop) CancelCurrentTask() bool { + al.currentCancelMu.Lock() + defer al.currentCancelMu.Unlock() + + if al.currentCancel == nil { + return false + } + + al.currentCancel() + al.currentCancel = nil + return true +} + func (al *AgentLoop) RegisterTool(tool tools.Tool) { for _, agentID := range al.registry.ListAgentIDs() { if agent, ok := al.registry.GetAgent(agentID); ok { @@ -781,6 +828,23 @@ func (al *AgentLoop) runAgentLoop( } } + // 0.5. Sanitize user message for sensitive data + sanitizedMsg := opts.UserMessage + var sanitizerMappings map[string]string + if sanitizer := utils.GetGlobalSanitizer(); sanitizer != nil { + result := sanitizer.Sanitize(opts.UserMessage) + sanitizedMsg = result.Sanitized + sanitizerMappings = result.Mappings + if len(sanitizerMappings) > 0 { + logger.DebugCF("agent", "Sanitized user message", + map[string]any{ + "original_len": len(opts.UserMessage), + "sanitized_len": len(sanitizedMsg), + "mappings": len(sanitizerMappings), + }) + } + } + // 1. Build messages (skip history for heartbeat) var history []providers.Message var summary string @@ -791,7 +855,7 @@ func (al *AgentLoop) runAgentLoop( messages := agent.ContextBuilder.BuildMessages( history, summary, - opts.UserMessage, + sanitizedMsg, // Use sanitized message for LLM opts.Media, opts.Channel, opts.ChatID, @@ -830,6 +894,17 @@ func (al *AgentLoop) runAgentLoop( finalContent = opts.DefaultResponse } + // 4.5. Restore sanitized content in LLM response + if len(sanitizerMappings) > 0 { + if sanitizer := utils.GetGlobalSanitizer(); sanitizer != nil { + finalContent = sanitizer.Restore(finalContent, sanitizerMappings) + logger.DebugCF("agent", "Restored sanitized content in response", + map[string]any{ + "mappings": len(sanitizerMappings), + }) + } + } + // 5. Save final assistant message to session agent.Sessions.AddMessage(opts.SessionKey, "assistant", finalContent) @@ -1655,6 +1730,43 @@ func (al *AgentLoop) estimateTokens(messages []providers.Message) int { return totalChars * 2 / 5 } +// handleCommandAsync handles commands in a separate goroutine so they can +// interrupt long-running tasks. Commands like /stop need to be processed +// immediately without waiting for the current message to finish. +func (al *AgentLoop) handleCommandAsync(ctx context.Context, msg bus.InboundMessage) { + logger.InfoCF("agent", "Processing command asynchronously", + map[string]any{ + "channel": msg.Channel, + "chat_id": msg.ChatID, + "sender_id": msg.SenderID, + "content": utils.Truncate(msg.Content, 50), + }) + + // Get default agent for command handling + agent := al.registry.GetDefaultAgent() + + response, handled := al.handleCommand(ctx, msg, agent) + if !handled { + // Command not recognized or passed through, ignore + return + } + + // Send response if any + if response != "" { + pubCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + if err := al.bus.PublishOutbound(pubCtx, bus.OutboundMessage{ + Channel: msg.Channel, + ChatID: msg.ChatID, + Content: response, + }); err != nil { + logger.WarnCF("agent", "Failed to send command response", + map[string]any{"error": err.Error()}) + } + } +} + func (al *AgentLoop) handleCommand( ctx context.Context, msg bus.InboundMessage, @@ -1717,6 +1829,7 @@ func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance) *commands.Runtim } return nil }, + CancelCurrentTask: al.CancelCurrentTask, } if agent != nil { rt.GetModelInfo = func() (string, string) { @@ -1770,3 +1883,34 @@ func extractParentPeer(msg bus.InboundMessage) *routing.RoutePeer { } return &routing.RoutePeer{Kind: parentKind, ID: parentID} } + +// convertSanitizerKeywords converts config keywords to utils keywords +func convertSanitizerKeywords(keywords []config.SanitizerKeyword) []utils.KeywordRule { + if len(keywords) == 0 { + return nil + } + result := make([]utils.KeywordRule, len(keywords)) + for i, kw := range keywords { + result[i] = utils.KeywordRule{ + Word: kw.Word, + Tag: kw.Tag, + } + } + return result +} + +// convertSanitizerPatterns converts config patterns to utils patterns +func convertSanitizerPatterns(patterns []config.SanitizerPattern) []utils.CustomPatternRule { + if len(patterns) == 0 { + return nil + } + result := make([]utils.CustomPatternRule, len(patterns)) + for i, p := range patterns { + result[i] = utils.CustomPatternRule{ + Name: p.Name, + Pattern: p.Pattern, + Tag: p.Tag, + } + } + return result +}