diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index 8bda60c..1c1b48f 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -477,6 +477,17 @@ func runServe(cmd *cobra.Command, args []string) error { slog.Info("message retention disabled") } + // Start stalemate worker (message acknowledgment enforcement) + stalemateConfig := messaging.ParseStalemateConfig() + stalemateWorker := messaging.NewStalemateWorker(db.DB, msgService, &channelLookupAdapter{channelService: channelService}, stalemateConfig) + stalemateWorker.Start() + slog.Info("stalemate worker started", + "processing_timeout", stalemateConfig.ProcessingTimeout.String(), + "reminder_after", stalemateConfig.ReminderAfter.String(), + "escalate_after", stalemateConfig.EscalateAfter.String(), + "interval", stalemateConfig.Interval.String(), + ) + // Create health checker healthChecker := health.NewChecker(db.DB, version) @@ -648,6 +659,9 @@ func runServe(cmd *cobra.Command, args []string) error { retentionWorker.Stop() } + // Stop stalemate worker + stalemateWorker.Stop() + // Stop embedding pipeline if embPipeline != nil { embPipeline.Stop() @@ -810,3 +824,16 @@ func ensureDefaultMCPClient(ctx context.Context, db *sql.DB, bcryptCost int) { "scopes", "mcp", ) } + +// channelLookupAdapter adapts channels.Service to messaging.ChannelLookup. +type channelLookupAdapter struct { + channelService *channels.Service +} + +func (a *channelLookupAdapter) GetChannelIDByName(ctx context.Context, name string) (int64, error) { + ch, err := a.channelService.GetChannelByName(ctx, name) + if err != nil { + return 0, err + } + return ch.ID, nil +} diff --git a/internal/messaging/stalemate.go b/internal/messaging/stalemate.go new file mode 100644 index 0000000..1e21a50 --- /dev/null +++ b/internal/messaging/stalemate.go @@ -0,0 +1,471 @@ +package messaging + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "log/slog" + "os" + "strconv" + "strings" + "sync" + "time" +) + +// StalemateConfig holds stalemate detection settings. +type StalemateConfig struct { + // ProcessingTimeout is how long a message can stay in "processing" before auto-fail (default 24h). + ProcessingTimeout time.Duration + // ReminderAfter is how long a pending DM waits before a system reminder is sent (default 4h). + ReminderAfter time.Duration + // EscalateAfter is how long a pending DM waits before escalation to #approvals (default 48h). + EscalateAfter time.Duration + // Interval is how often the worker checks for stale messages (default 15m). + Interval time.Duration +} + +// DefaultStalemateConfig returns the default stalemate configuration. +func DefaultStalemateConfig() StalemateConfig { + return StalemateConfig{ + ProcessingTimeout: 24 * time.Hour, + ReminderAfter: 4 * time.Hour, + EscalateAfter: 48 * time.Hour, + Interval: 15 * time.Minute, + } +} + +// parseDurationWithDays parses a duration string supporting "Nd" format for days +// in addition to standard Go duration formats. +func parseDurationWithDays(s string) (time.Duration, error) { + s = strings.TrimSpace(s) + if s == "" { + return 0, fmt.Errorf("empty duration string") + } + + // Try "Nd" format (days) + if strings.HasSuffix(s, "d") { + days, err := strconv.Atoi(strings.TrimSuffix(s, "d")) + if err == nil && days > 0 { + return time.Duration(days) * 24 * time.Hour, nil + } + } + + // Try standard Go duration + return time.ParseDuration(s) +} + +// ParseStalemateConfig reads stalemate configuration from environment variables. +func ParseStalemateConfig() StalemateConfig { + cfg := DefaultStalemateConfig() + + if v := os.Getenv("SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT"); v != "" { + if d, err := parseDurationWithDays(v); err == nil && d > 0 { + cfg.ProcessingTimeout = d + } + } + if v := os.Getenv("SYNAPBUS_STALEMATE_REMINDER_AFTER"); v != "" { + if d, err := parseDurationWithDays(v); err == nil && d > 0 { + cfg.ReminderAfter = d + } + } + if v := os.Getenv("SYNAPBUS_STALEMATE_ESCALATE_AFTER"); v != "" { + if d, err := parseDurationWithDays(v); err == nil && d > 0 { + cfg.EscalateAfter = d + } + } + if v := os.Getenv("SYNAPBUS_STALEMATE_INTERVAL"); v != "" { + if d, err := parseDurationWithDays(v); err == nil && d > 0 { + cfg.Interval = d + } + } + + return cfg +} + +// ChannelLookup provides channel lookup by name without importing the channels package. +type ChannelLookup interface { + // GetChannelIDByName returns a channel ID by name, or 0 if not found. + GetChannelIDByName(ctx context.Context, name string) (int64, error) +} + +// StalemateWorker periodically checks for and handles stale messages. +type StalemateWorker struct { + db *sql.DB + msgService *MessagingService + channelLookup ChannelLookup + config StalemateConfig + logger *slog.Logger + done chan struct{} + wg sync.WaitGroup +} + +// NewStalemateWorker creates a new stalemate detection worker. +func NewStalemateWorker(db *sql.DB, msgService *MessagingService, channelLookup ChannelLookup, config StalemateConfig) *StalemateWorker { + return &StalemateWorker{ + db: db, + msgService: msgService, + channelLookup: channelLookup, + config: config, + logger: slog.Default().With("component", "stalemate-worker"), + done: make(chan struct{}), + } +} + +// Start begins the background stalemate check loop. +func (w *StalemateWorker) Start() { + w.wg.Add(1) + go func() { + defer w.wg.Done() + w.logger.Info("stalemate worker started", + "interval", w.config.Interval.String(), + "processing_timeout", w.config.ProcessingTimeout.String(), + "reminder_after", w.config.ReminderAfter.String(), + "escalate_after", w.config.EscalateAfter.String(), + ) + + ticker := time.NewTicker(w.config.Interval) + defer ticker.Stop() + + for { + select { + case <-ticker.C: + ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) + w.checkStaleMessages(ctx) + cancel() + case <-w.done: + w.logger.Info("stalemate worker stopped") + return + } + } + }() +} + +// Stop stops the stalemate worker and waits for it to finish. +func (w *StalemateWorker) Stop() { + close(w.done) + w.wg.Wait() +} + +// checkStaleMessages runs all stalemate checks. +func (w *StalemateWorker) checkStaleMessages(ctx context.Context) { + failed := w.failTimedOutProcessing(ctx) + reminded := w.sendPendingReminders(ctx) + escalated := w.escalatePendingMessages(ctx) + + if failed > 0 || reminded > 0 || escalated > 0 { + w.logger.Info("stalemate check complete", + "auto_failed", failed, + "reminders_sent", reminded, + "escalations_sent", escalated, + ) + } +} + +// staleDM represents a stale direct message found by the worker. +type staleDM struct { + ID int64 + FromAgent string + ToAgent string + Body string + ClaimedAt *time.Time + ClaimedBy string + CreatedAt time.Time +} + +// failTimedOutProcessing auto-fails DMs in "processing" status that have exceeded the timeout. +func (w *StalemateWorker) failTimedOutProcessing(ctx context.Context) int64 { + cutoff := time.Now().Add(-w.config.ProcessingTimeout) + + rows, err := w.db.QueryContext(ctx, + `SELECT id, from_agent, to_agent, body, claimed_at, claimed_by + FROM messages + WHERE status = 'processing' + AND to_agent IS NOT NULL + AND to_agent != '' + AND to_agent != 'system' + AND claimed_at < ?`, + cutoff, + ) + if err != nil { + w.logger.Error("query timed-out processing messages failed", "error", err) + return 0 + } + defer rows.Close() + + var stale []staleDM + for rows.Next() { + var dm staleDM + var claimedAt sql.NullTime + var claimedBy sql.NullString + if err := rows.Scan(&dm.ID, &dm.FromAgent, &dm.ToAgent, &dm.Body, &claimedAt, &claimedBy); err != nil { + w.logger.Error("scan timed-out message failed", "error", err) + continue + } + if claimedAt.Valid { + dm.ClaimedAt = &claimedAt.Time + } + if claimedBy.Valid { + dm.ClaimedBy = claimedBy.String + } + stale = append(stale, dm) + } + + count := int64(0) + for _, dm := range stale { + metadata := map[string]any{"error": "claim timeout exceeded"} + metaBytes, _ := json.Marshal(metadata) + + // Update directly via DB since the store's UpdateMessageStatus requires the claiming agent + _, err := w.db.ExecContext(ctx, + `UPDATE messages SET status = ?, metadata = ?, updated_at = CURRENT_TIMESTAMP + WHERE id = ? AND status = 'processing'`, + StatusFailed, string(metaBytes), dm.ID, + ) + if err != nil { + w.logger.Error("auto-fail message failed", + "message_id", dm.ID, + "error", err, + ) + continue + } + w.logger.Info("auto-failed stale processing message", + "message_id", dm.ID, + "from_agent", dm.FromAgent, + "to_agent", dm.ToAgent, + "claimed_by", dm.ClaimedBy, + ) + count++ + } + return count +} + +// sendPendingReminders sends system DM reminders for pending messages older than ReminderAfter. +func (w *StalemateWorker) sendPendingReminders(ctx context.Context) int64 { + cutoff := time.Now().Add(-w.config.ReminderAfter) + + rows, err := w.db.QueryContext(ctx, + `SELECT id, from_agent, to_agent, body, created_at + FROM messages + WHERE status = 'pending' + AND to_agent IS NOT NULL + AND to_agent != '' + AND from_agent != 'system' + AND to_agent != 'system' + AND created_at < ?`, + cutoff, + ) + if err != nil { + w.logger.Error("query pending reminder candidates failed", "error", err) + return 0 + } + defer rows.Close() + + type pendingMsg struct { + ID int64 + FromAgent string + ToAgent string + Body string + CreatedAt time.Time + } + + var pending []pendingMsg + for rows.Next() { + var pm pendingMsg + if err := rows.Scan(&pm.ID, &pm.FromAgent, &pm.ToAgent, &pm.Body, &pm.CreatedAt); err != nil { + w.logger.Error("scan pending message failed", "error", err) + continue + } + pending = append(pending, pm) + } + + count := int64(0) + for _, pm := range pending { + // Check if a reminder already exists for this message + if w.reminderExists(ctx, pm.ID, pm.ToAgent) { + continue + } + + age := formatAge(time.Since(pm.CreatedAt)) + truncBody := truncate(pm.Body, 100) + + body := fmt.Sprintf( + "**Reminder**: You have a pending message from %s (%s old). Message: \"%s\". Please claim and process it.", + pm.FromAgent, age, truncBody, + ) + + _, err := w.msgService.SendMessage(ctx, "system", pm.ToAgent, body, SendOptions{ + Subject: fmt.Sprintf("stalemate-reminder:%d", pm.ID), + Priority: 7, + Metadata: fmt.Sprintf(`{"stalemate_reminder_for":%d}`, pm.ID), + }) + if err != nil { + w.logger.Error("send stalemate reminder failed", + "message_id", pm.ID, + "to_agent", pm.ToAgent, + "error", err, + ) + continue + } + w.logger.Info("sent stalemate reminder", + "message_id", pm.ID, + "to_agent", pm.ToAgent, + "from_agent", pm.FromAgent, + "age", age, + ) + count++ + } + return count +} + +// escalatePendingMessages escalates pending messages older than EscalateAfter to #approvals. +func (w *StalemateWorker) escalatePendingMessages(ctx context.Context) int64 { + cutoff := time.Now().Add(-w.config.EscalateAfter) + + rows, err := w.db.QueryContext(ctx, + `SELECT id, from_agent, to_agent, body, created_at + FROM messages + WHERE status = 'pending' + AND to_agent IS NOT NULL + AND to_agent != '' + AND from_agent != 'system' + AND to_agent != 'system' + AND created_at < ?`, + cutoff, + ) + if err != nil { + w.logger.Error("query escalation candidates failed", "error", err) + return 0 + } + defer rows.Close() + + type pendingMsg struct { + ID int64 + FromAgent string + ToAgent string + Body string + CreatedAt time.Time + } + + var pending []pendingMsg + for rows.Next() { + var pm pendingMsg + if err := rows.Scan(&pm.ID, &pm.FromAgent, &pm.ToAgent, &pm.Body, &pm.CreatedAt); err != nil { + w.logger.Error("scan escalation candidate failed", "error", err) + continue + } + pending = append(pending, pm) + } + + if len(pending) == 0 { + return 0 + } + + // Look up #approvals channel + channelID, err := w.channelLookup.GetChannelIDByName(ctx, "approvals") + if err != nil { + w.logger.Warn("cannot escalate: #approvals channel not found", "error", err) + return 0 + } + + count := int64(0) + for _, pm := range pending { + // Check if already escalated + if w.escalationExists(ctx, pm.ID) { + continue + } + + age := formatAge(time.Since(pm.CreatedAt)) + truncBody := truncate(pm.Body, 100) + + body := fmt.Sprintf( + "**ESCALATION**: Pending message for @%s from %s has been unprocessed for %s. Message: \"%s\". Manual intervention may be required.", + pm.ToAgent, pm.FromAgent, age, truncBody, + ) + + _, err := w.msgService.SendMessage(ctx, "system", "", body, SendOptions{ + Subject: fmt.Sprintf("stalemate-escalation:%d", pm.ID), + Priority: 9, + Metadata: fmt.Sprintf(`{"stalemate_escalation_for":%d}`, pm.ID), + ChannelID: &channelID, + }) + if err != nil { + w.logger.Error("send escalation to #approvals failed", + "message_id", pm.ID, + "error", err, + ) + continue + } + w.logger.Info("escalated stale message to #approvals", + "message_id", pm.ID, + "to_agent", pm.ToAgent, + "from_agent", pm.FromAgent, + "age", age, + ) + count++ + } + return count +} + +// reminderExists checks if a system reminder already exists for a given message ID. +func (w *StalemateWorker) reminderExists(ctx context.Context, messageID int64, toAgent string) bool { + var count int + err := w.db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages + WHERE from_agent = 'system' + AND to_agent = ? + AND metadata LIKE ?`, + toAgent, fmt.Sprintf(`%%"stalemate_reminder_for":%d%%`, messageID), + ).Scan(&count) + if err != nil { + return false + } + return count > 0 +} + +// escalationExists checks if an escalation already exists for a given message ID. +func (w *StalemateWorker) escalationExists(ctx context.Context, messageID int64) bool { + var count int + err := w.db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages + WHERE from_agent = 'system' + AND metadata LIKE ?`, + fmt.Sprintf(`%%"stalemate_escalation_for":%d%%`, messageID), + ).Scan(&count) + if err != nil { + return false + } + return count > 0 +} + +// truncate truncates a string to maxLen characters, appending "..." if truncated. +func truncate(s string, maxLen int) string { + runes := []rune(s) + if len(runes) <= maxLen { + return s + } + return string(runes[:maxLen]) + "..." +} + +// formatAge returns a human-readable age string. +func formatAge(d time.Duration) string { + if d < time.Hour { + return fmt.Sprintf("%dm", int(d.Minutes())) + } + hours := int(d.Hours()) + if hours < 24 { + return fmt.Sprintf("%dh", hours) + } + days := hours / 24 + remainingHours := hours % 24 + if remainingHours == 0 { + if days == 1 { + return "1 day" + } + return fmt.Sprintf("%d days", days) + } + if days == 1 { + return fmt.Sprintf("1 day %dh", remainingHours) + } + return fmt.Sprintf("%d days %dh", days, remainingHours) +} diff --git a/internal/messaging/stalemate_test.go b/internal/messaging/stalemate_test.go new file mode 100644 index 0000000..e0035ed --- /dev/null +++ b/internal/messaging/stalemate_test.go @@ -0,0 +1,480 @@ +package messaging + +import ( + "context" + "database/sql" + "fmt" + "os" + "testing" + "time" + + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/trace" +) + +// stubChannelLookup implements ChannelLookup for tests. +type stubChannelLookup struct { + channelID int64 + err error +} + +func (s *stubChannelLookup) GetChannelIDByName(ctx context.Context, name string) (int64, error) { + if s.err != nil { + return 0, s.err + } + return s.channelID, nil +} + +// newStalemateTestService creates a MessagingService and DB for stalemate tests. +func newStalemateTestService(t *testing.T) (*MessagingService, *sql.DB) { + t.Helper() + db := newTestDB(t) + + seedAgent(t, db, "sender") + seedAgent(t, db, "receiver") + seedAgent(t, db, "system") + + store := NewSQLiteMessageStore(db) + tracer := trace.NewTracer(db) + t.Cleanup(func() { tracer.Close() }) + + svc := NewMessagingService(store, tracer) + return svc, db +} + +// insertStaleMessage inserts a message with a specific created_at and claimed_at for testing. +func insertStaleMessage(t *testing.T, db *sql.DB, from, to, body, status string, createdAt time.Time, claimedAt *time.Time, claimedBy string) int64 { + t.Helper() + + // Insert conversation first + result, err := db.Exec( + `INSERT INTO conversations (subject, created_by, created_at, updated_at) + VALUES (?, ?, ?, ?)`, + "stalemate-test", from, createdAt, createdAt, + ) + if err != nil { + t.Fatalf("insert conversation: %v", err) + } + convID, _ := result.LastInsertId() + + var claimedAtSQL interface{} = nil + if claimedAt != nil { + claimedAtSQL = *claimedAt + } + var claimedBySQL interface{} = nil + if claimedBy != "" { + claimedBySQL = claimedBy + } + + result, err = db.Exec( + `INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, claimed_by, claimed_at, created_at, updated_at) + VALUES (?, ?, ?, ?, 5, ?, '{}', ?, ?, ?, ?)`, + convID, from, to, body, status, claimedBySQL, claimedAtSQL, createdAt, createdAt, + ) + if err != nil { + t.Fatalf("insert stale message: %v", err) + } + id, _ := result.LastInsertId() + return id +} + +func TestStalemateWorker_ProcessingTimeout(t *testing.T) { + svc, db := newStalemateTestService(t) + ctx := context.Background() + + // Insert a message in "processing" status with old claimed_at + oldClaimedAt := time.Now().Add(-25 * time.Hour) + msgID := insertStaleMessage(t, db, "sender", "receiver", "stale processing task", StatusProcessing, time.Now().Add(-26*time.Hour), &oldClaimedAt, "receiver") + + config := DefaultStalemateConfig() + config.ProcessingTimeout = 24 * time.Hour + + lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} + worker := NewStalemateWorker(db, svc, lookup, config) + + worker.checkStaleMessages(ctx) + + // Verify message was auto-failed + var status, metadata string + err := db.QueryRowContext(ctx, `SELECT status, metadata FROM messages WHERE id = ?`, msgID).Scan(&status, &metadata) + if err != nil { + t.Fatalf("query message: %v", err) + } + if status != StatusFailed { + t.Errorf("status = %q, want %q", status, StatusFailed) + } + if metadata == "{}" { + t.Error("expected metadata to contain error info") + } +} + +func TestStalemateWorker_ProcessingTimeout_NotExpired(t *testing.T) { + svc, db := newStalemateTestService(t) + ctx := context.Background() + + // Insert a message in "processing" status with recent claimed_at (should NOT be failed) + recentClaimedAt := time.Now().Add(-1 * time.Hour) + msgID := insertStaleMessage(t, db, "sender", "receiver", "recent processing task", StatusProcessing, time.Now().Add(-2*time.Hour), &recentClaimedAt, "receiver") + + config := DefaultStalemateConfig() + config.ProcessingTimeout = 24 * time.Hour + + lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} + worker := NewStalemateWorker(db, svc, lookup, config) + + worker.checkStaleMessages(ctx) + + // Verify message was NOT auto-failed + var status string + err := db.QueryRowContext(ctx, `SELECT status FROM messages WHERE id = ?`, msgID).Scan(&status) + if err != nil { + t.Fatalf("query message: %v", err) + } + if status != StatusProcessing { + t.Errorf("status = %q, want %q (should not have been failed)", status, StatusProcessing) + } +} + +func TestStalemateWorker_PendingReminder(t *testing.T) { + svc, db := newStalemateTestService(t) + ctx := context.Background() + + // Insert a pending DM that is 5 hours old + insertStaleMessage(t, db, "sender", "receiver", "please review this", StatusPending, time.Now().Add(-5*time.Hour), nil, "") + + config := DefaultStalemateConfig() + config.ReminderAfter = 4 * time.Hour + config.EscalateAfter = 48 * time.Hour // won't trigger + + lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} + worker := NewStalemateWorker(db, svc, lookup, config) + + worker.checkStaleMessages(ctx) + + // Verify a system reminder was sent to receiver + var count int + err := db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND to_agent = 'receiver' AND body LIKE '%Reminder%'`, + ).Scan(&count) + if err != nil { + t.Fatalf("query reminder: %v", err) + } + if count != 1 { + t.Errorf("expected 1 reminder, got %d", count) + } +} + +func TestStalemateWorker_SystemMessageSkip(t *testing.T) { + svc, db := newStalemateTestService(t) + ctx := context.Background() + + // Insert a pending DM FROM system (should be skipped) + insertStaleMessage(t, db, "system", "receiver", "system notification", StatusPending, time.Now().Add(-5*time.Hour), nil, "") + + config := DefaultStalemateConfig() + config.ReminderAfter = 4 * time.Hour + + lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} + worker := NewStalemateWorker(db, svc, lookup, config) + + worker.checkStaleMessages(ctx) + + // Verify NO reminder was sent (only the original system message should exist) + var count int + err := db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND body LIKE '%Reminder%'`, + ).Scan(&count) + if err != nil { + t.Fatalf("query reminder: %v", err) + } + if count != 0 { + t.Errorf("expected 0 reminders for system message, got %d", count) + } +} + +func TestStalemateWorker_DuplicateReminderPrevention(t *testing.T) { + svc, db := newStalemateTestService(t) + ctx := context.Background() + + // Insert a pending DM that is old enough for a reminder + insertStaleMessage(t, db, "sender", "receiver", "need your attention", StatusPending, time.Now().Add(-5*time.Hour), nil, "") + + config := DefaultStalemateConfig() + config.ReminderAfter = 4 * time.Hour + config.EscalateAfter = 48 * time.Hour + + lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} + worker := NewStalemateWorker(db, svc, lookup, config) + + // Run check twice + worker.checkStaleMessages(ctx) + worker.checkStaleMessages(ctx) + + // Verify only ONE reminder was sent + var count int + err := db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND to_agent = 'receiver' AND body LIKE '%Reminder%'`, + ).Scan(&count) + if err != nil { + t.Fatalf("query reminders: %v", err) + } + if count != 1 { + t.Errorf("expected 1 reminder (no duplicates), got %d", count) + } +} + +func TestStalemateWorker_Escalation(t *testing.T) { + svc, db := newStalemateTestService(t) + ctx := context.Background() + + // Create #approvals channel + _, err := db.Exec( + `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at) + VALUES (1, 'approvals', 'Approval queue', '', 'standard', 0, 0, 'system', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) + if err != nil { + t.Fatalf("create approvals channel: %v", err) + } + // Add system as member + _, err = db.Exec( + `INSERT INTO channel_members (channel_id, agent_name, role, joined_at) + VALUES (1, 'system', 'owner', CURRENT_TIMESTAMP)`) + if err != nil { + t.Fatalf("add system to channel: %v", err) + } + + // Insert a pending DM that is 49 hours old (beyond escalation threshold) + insertStaleMessage(t, db, "sender", "receiver", "urgent task ignored", StatusPending, time.Now().Add(-49*time.Hour), nil, "") + + config := DefaultStalemateConfig() + config.ReminderAfter = 4 * time.Hour + config.EscalateAfter = 48 * time.Hour + + lookup := &stubChannelLookup{channelID: 1} + worker := NewStalemateWorker(db, svc, lookup, config) + + worker.checkStaleMessages(ctx) + + // Verify an escalation was sent to #approvals channel + var count int + err = db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND channel_id = 1 AND body LIKE '%ESCALATION%'`, + ).Scan(&count) + if err != nil { + t.Fatalf("query escalations: %v", err) + } + if count != 1 { + t.Errorf("expected 1 escalation, got %d", count) + } +} + +func TestStalemateWorker_DuplicateEscalationPrevention(t *testing.T) { + svc, db := newStalemateTestService(t) + ctx := context.Background() + + // Create #approvals channel + db.Exec( + `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at) + VALUES (1, 'approvals', 'Approval queue', '', 'standard', 0, 0, 'system', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) + db.Exec( + `INSERT INTO channel_members (channel_id, agent_name, role, joined_at) + VALUES (1, 'system', 'owner', CURRENT_TIMESTAMP)`) + + // Insert a pending DM that is 49 hours old + insertStaleMessage(t, db, "sender", "receiver", "urgent task", StatusPending, time.Now().Add(-49*time.Hour), nil, "") + + config := DefaultStalemateConfig() + config.ReminderAfter = 4 * time.Hour + config.EscalateAfter = 48 * time.Hour + + lookup := &stubChannelLookup{channelID: 1} + worker := NewStalemateWorker(db, svc, lookup, config) + + // Run check twice + worker.checkStaleMessages(ctx) + worker.checkStaleMessages(ctx) + + // Verify only ONE escalation was sent + var count int + err := db.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND channel_id = 1 AND body LIKE '%ESCALATION%'`, + ).Scan(&count) + if err != nil { + t.Fatalf("query escalations: %v", err) + } + if count != 1 { + t.Errorf("expected 1 escalation (no duplicates), got %d", count) + } +} + +func TestParseStalemateConfig(t *testing.T) { + tests := []struct { + name string + envVars map[string]string + expected StalemateConfig + }{ + { + name: "defaults when no env vars", + envVars: map[string]string{}, + expected: DefaultStalemateConfig(), + }, + { + name: "custom values with day format", + envVars: map[string]string{ + "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "7d", + "SYNAPBUS_STALEMATE_REMINDER_AFTER": "8h", + "SYNAPBUS_STALEMATE_ESCALATE_AFTER": "3d", + "SYNAPBUS_STALEMATE_INTERVAL": "30m", + }, + expected: StalemateConfig{ + ProcessingTimeout: 7 * 24 * time.Hour, + ReminderAfter: 8 * time.Hour, + EscalateAfter: 3 * 24 * time.Hour, + Interval: 30 * time.Minute, + }, + }, + { + name: "standard Go duration format", + envVars: map[string]string{ + "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "48h", + "SYNAPBUS_STALEMATE_REMINDER_AFTER": "2h30m", + "SYNAPBUS_STALEMATE_ESCALATE_AFTER": "72h", + "SYNAPBUS_STALEMATE_INTERVAL": "5m", + }, + expected: StalemateConfig{ + ProcessingTimeout: 48 * time.Hour, + ReminderAfter: 2*time.Hour + 30*time.Minute, + EscalateAfter: 72 * time.Hour, + Interval: 5 * time.Minute, + }, + }, + { + name: "invalid values fall back to defaults", + envVars: map[string]string{ + "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "invalid", + "SYNAPBUS_STALEMATE_REMINDER_AFTER": "bad", + "SYNAPBUS_STALEMATE_ESCALATE_AFTER": "", + "SYNAPBUS_STALEMATE_INTERVAL": "-5m", + }, + expected: DefaultStalemateConfig(), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Clear all env vars first + envKeys := []string{ + "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT", + "SYNAPBUS_STALEMATE_REMINDER_AFTER", + "SYNAPBUS_STALEMATE_ESCALATE_AFTER", + "SYNAPBUS_STALEMATE_INTERVAL", + } + for _, k := range envKeys { + os.Unsetenv(k) + } + + // Set test env vars + for k, v := range tt.envVars { + os.Setenv(k, v) + } + defer func() { + for _, k := range envKeys { + os.Unsetenv(k) + } + }() + + cfg := ParseStalemateConfig() + + if cfg.ProcessingTimeout != tt.expected.ProcessingTimeout { + t.Errorf("ProcessingTimeout = %v, want %v", cfg.ProcessingTimeout, tt.expected.ProcessingTimeout) + } + if cfg.ReminderAfter != tt.expected.ReminderAfter { + t.Errorf("ReminderAfter = %v, want %v", cfg.ReminderAfter, tt.expected.ReminderAfter) + } + if cfg.EscalateAfter != tt.expected.EscalateAfter { + t.Errorf("EscalateAfter = %v, want %v", cfg.EscalateAfter, tt.expected.EscalateAfter) + } + if cfg.Interval != tt.expected.Interval { + t.Errorf("Interval = %v, want %v", cfg.Interval, tt.expected.Interval) + } + }) + } +} + +func TestParseDurationWithDays(t *testing.T) { + tests := []struct { + name string + input string + want time.Duration + wantErr bool + }{ + {"7 days", "7d", 7 * 24 * time.Hour, false}, + {"1 day", "1d", 24 * time.Hour, false}, + {"30 days", "30d", 30 * 24 * time.Hour, false}, + {"standard hours", "48h", 48 * time.Hour, false}, + {"standard minutes", "15m", 15 * time.Minute, false}, + {"mixed duration", "2h30m", 2*time.Hour + 30*time.Minute, false}, + {"empty string", "", 0, true}, + {"invalid", "xyz", 0, true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := parseDurationWithDays(tt.input) + if (err != nil) != tt.wantErr { + t.Errorf("parseDurationWithDays(%q) error = %v, wantErr %v", tt.input, err, tt.wantErr) + return + } + if got != tt.want { + t.Errorf("parseDurationWithDays(%q) = %v, want %v", tt.input, got, tt.want) + } + }) + } +} + +func TestTruncate(t *testing.T) { + tests := []struct { + name string + input string + maxLen int + want string + }{ + {"short string", "hello", 10, "hello"}, + {"exact length", "hello", 5, "hello"}, + {"truncated", "hello world, this is a long message", 10, "hello worl..."}, + {"empty", "", 10, ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := truncate(tt.input, tt.maxLen) + if got != tt.want { + t.Errorf("truncate(%q, %d) = %q, want %q", tt.input, tt.maxLen, got, tt.want) + } + }) + } +} + +func TestFormatAge(t *testing.T) { + tests := []struct { + name string + d time.Duration + want string + }{ + {"minutes", 30 * time.Minute, "30m"}, + {"hours", 5 * time.Hour, "5h"}, + {"1 day", 24 * time.Hour, "1 day"}, + {"2 days", 48 * time.Hour, "2 days"}, + {"1 day with hours", 25 * time.Hour, "1 day 1h"}, + {"2 days with hours", 50 * time.Hour, "2 days 2h"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := formatAge(tt.d) + if got != tt.want { + t.Errorf("formatAge(%v) = %q, want %q", tt.d, got, tt.want) + } + }) + } +}