diff --git a/pkg/agent/interruption_checker.go b/pkg/agent/interruption_checker.go index 0d6d02cfd..ffb128761 100644 --- a/pkg/agent/interruption_checker.go +++ b/pkg/agent/interruption_checker.go @@ -7,6 +7,7 @@ package agent import ( "sync" + "time" "github.com/sipeed/picoclaw/pkg/bus" ) @@ -20,24 +21,36 @@ import ( // - Thread-safe for concurrent access // - Simple API: Signal, DrainAll, HasPending // - Zero overhead when not in use +// - Grace period to handle race conditions type InterruptionChecker struct { - queue []bus.InboundMessage - mu sync.Mutex + queue []bus.InboundMessage + mu sync.Mutex + gracePeriodEnd time.Time // Allow signaling even after session "ends" + active bool // Whether session is actively processing } // NewInterruptionChecker creates a new checker for a session func NewInterruptionChecker() *InterruptionChecker { return &InterruptionChecker{ - queue: make([]bus.InboundMessage, 0, 10), // Pre-allocate for common case + queue: make([]bus.InboundMessage, 0, 10), // Pre-allocate for common case + active: true, // Start as active } } // Signal pushes a new interrupting message into the queue. // This is called when a new message arrives for an already-active session. -func (ic *InterruptionChecker) Signal(msg bus.InboundMessage) { +// Messages are accepted during active period OR within grace period. +func (ic *InterruptionChecker) Signal(msg bus.InboundMessage) bool { ic.mu.Lock() defer ic.mu.Unlock() - ic.queue = append(ic.queue, msg) + + // Accept message if session is active OR within grace period + if ic.active || time.Now().Before(ic.gracePeriodEnd) { + ic.queue = append(ic.queue, msg) + return true + } + + return false } // DrainAll returns and clears all pending messages. @@ -93,3 +106,27 @@ func (ic *InterruptionChecker) Clear() { defer ic.mu.Unlock() ic.queue = ic.queue[:0] } + +// SetGracePeriod sets a grace period during which messages can still be signaled +// even after the session becomes "inactive". This handles race conditions where +// messages arrive just as the LLM is finishing. +func (ic *InterruptionChecker) SetGracePeriod(duration time.Duration) { + ic.mu.Lock() + defer ic.mu.Unlock() + ic.gracePeriodEnd = time.Now().Add(duration) + ic.active = false // Mark as inactive but accept messages during grace period +} + +// IsActive returns true if the checker is still active or within grace period +func (ic *InterruptionChecker) IsActive() bool { + ic.mu.Lock() + defer ic.mu.Unlock() + return ic.active || time.Now().Before(ic.gracePeriodEnd) +} + +// Deactivate marks the checker as inactive (but respects grace period) +func (ic *InterruptionChecker) Deactivate() { + ic.mu.Lock() + defer ic.mu.Unlock() + ic.active = false +} diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index 59a7d3ae7..bc231b52f 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -232,16 +232,27 @@ func (al *AgentLoop) Run(ctx context.Context) error { if al.hasActiveChecker(sessionKey) { // Session is active, signal interruption instead of creating new task checker := al.getOrCreateChecker(sessionKey) - checker.Signal(msg) + signaled := checker.Signal(msg) - logger.InfoCF("agent", "Steering: signaled interruption for active session", - map[string]any{ - "session_key": sessionKey, - "channel": msg.Channel, - "chat_id": msg.ChatID, - "content_preview": utils.Truncate(msg.Content, 60), - }) - continue // Don't process as new message + if signaled { + logger.InfoCF("agent", "Steering: signaled interruption for active session", + map[string]any{ + "session_key": sessionKey, + "channel": msg.Channel, + "chat_id": msg.ChatID, + "content_preview": utils.Truncate(msg.Content, 60), + }) + continue // Don't process as new message + } else { + // Grace period expired, treat as new message + logger.WarnCF("agent", "Steering: session checker exists but grace period expired, processing as new message", + map[string]any{ + "session_key": sessionKey, + "channel": msg.Channel, + "chat_id": msg.ChatID, + }) + // Fall through to process as new message + } } } @@ -924,10 +935,42 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt // ===== NEW: Steering Architecture - Setup checker for this session ===== if al.enableSteering { // Create checker to signal this session is active - al.getOrCreateChecker(opts.SessionKey) + checker := al.getOrCreateChecker(opts.SessionKey) - // Cleanup checker when done - defer al.removeChecker(opts.SessionKey) + // Cleanup with grace period to handle race conditions + defer func() { + // Set grace period before checking for late arrivals + checker.SetGracePeriod(2 * time.Second) + + // Wait briefly for any race-condition messages + time.Sleep(150 * time.Millisecond) + + // Check one final time for pending interruptions + finalPending := checker.DrainAll() + if len(finalPending) > 0 { + logger.InfoCF("agent", "Steering: found interruptions during grace period, reprocessing", + map[string]any{ + "session_key": opts.SessionKey, + "pending_count": len(finalPending), + "channel": opts.Channel, + "chat_id": opts.ChatID, + }) + + // Re-trigger processing by publishing as new inbound message + // Combine all pending messages + injectionContent := formatInterruptionInjection(finalPending) + al.bus.PublishInbound(ctx, bus.InboundMessage{ + Channel: opts.Channel, + ChatID: opts.ChatID, + SessionKey: opts.SessionKey, + Content: injectionContent, + Metadata: make(map[string]string), // Empty metadata for re-triggered message + }) + } + + // Now safe to remove checker + al.removeChecker(opts.SessionKey) + }() } // 0a. Update interrupt handler context @@ -1271,7 +1314,7 @@ func (al *AgentLoop) runLLMIteration( if al.enableSteering { checker := al.getOrCreateChecker(opts.SessionKey) pending := checker.DrainAll() - + if len(pending) > 0 { logger.InfoCF("agent", "Steering: LLM finished but has pending interruptions, injecting", map[string]any{ @@ -1279,7 +1322,7 @@ func (al *AgentLoop) runLLMIteration( "pending_count": len(pending), "iteration": iteration, }) - + // Save the assistant's response first assistantMsg := providers.Message{ Role: "assistant", @@ -1287,7 +1330,7 @@ func (al *AgentLoop) runLLMIteration( } messages = append(messages, assistantMsg) agent.Sessions.AddMessage(opts.SessionKey, "assistant", response.Content) - + // Send the response to user if !constants.IsInternalChannel(opts.Channel) { al.bus.PublishOutbound(ctx, bus.OutboundMessage{ @@ -1296,7 +1339,7 @@ func (al *AgentLoop) runLLMIteration( Content: response.Content, }) } - + // Format and inject interruption injectionContent := formatInterruptionInjection(pending) injectionMsg := providers.Message{ @@ -1305,12 +1348,12 @@ func (al *AgentLoop) runLLMIteration( } messages = append(messages, injectionMsg) agent.Sessions.AddMessage(opts.SessionKey, "user", injectionContent) - + // Continue to handle the interruption continue } } - + // No interruptions, finish normally finalContent = response.Content logger.InfoCF("agent", "LLM response without tool calls (direct answer)", @@ -1319,6 +1362,53 @@ func (al *AgentLoop) runLLMIteration( "iteration": iteration, "content_chars": len(finalContent), }) + + // FINAL SAFETY CHECK: One more check for race-condition interruptions + // This catches messages that arrived while we were processing the final response + if al.enableSteering { + time.Sleep(100 * time.Millisecond) // Brief wait for any in-flight messages + checker := al.getOrCreateChecker(opts.SessionKey) + lastMinutePending := checker.DrainAll() + + if len(lastMinutePending) > 0 { + logger.InfoCF("agent", "Steering: caught last-minute interruptions after final response", + map[string]any{ + "session_key": opts.SessionKey, + "pending_count": len(lastMinutePending), + "iteration": iteration, + }) + + // Save assistant's response first + assistantMsg := providers.Message{ + Role: "assistant", + Content: response.Content, + } + messages = append(messages, assistantMsg) + agent.Sessions.AddMessage(opts.SessionKey, "assistant", response.Content) + + // Send response to user + if !constants.IsInternalChannel(opts.Channel) { + al.bus.PublishOutbound(ctx, bus.OutboundMessage{ + Channel: opts.Channel, + ChatID: opts.ChatID, + Content: response.Content, + }) + } + + // Inject last-minute interruptions + injectionContent := formatInterruptionInjection(lastMinutePending) + injectionMsg := providers.Message{ + Role: "user", + Content: injectionContent, + } + messages = append(messages, injectionMsg) + agent.Sessions.AddMessage(opts.SessionKey, "user", injectionContent) + + // Continue to handle these interruptions + continue + } + } + break }