diff --git a/gateway/platforms/helpers.py b/gateway/platforms/helpers.py index 1ac1d431c..85edb02c5 100644 --- a/gateway/platforms/helpers.py +++ b/gateway/platforms/helpers.py @@ -82,6 +82,10 @@ class MessageDeduplicator: del self._seen[msg_id] return False + def discard(self, msg_id: str) -> None: + """Release a claimed message ID after cancelled/failed handoff.""" + self._seen.pop(msg_id, None) + def clear(self): """Clear all tracked messages.""" self._seen.clear() diff --git a/plugins/platforms/discord/adapter.py b/plugins/platforms/discord/adapter.py index fe8784106..d8973b51f 100644 --- a/plugins/platforms/discord/adapter.py +++ b/plugins/platforms/discord/adapter.py @@ -18,7 +18,6 @@ import logging import math import os import re -import sqlite3 import struct import subprocess import tempfile @@ -52,8 +51,7 @@ _DISCORD_COMMAND_SYNC_POLICIES = {"safe", "bulk", "off"} _DISCORD_COMMAND_SYNC_STATE_SUBDIR = "gateway" _DISCORD_COMMAND_SYNC_STATE_FILENAME = "discord_command_sync_state.json" _DISCORD_NONCONVERSATIONAL_STATE_FILENAME = "discord_nonconversational_messages.json" -_DISCORD_RECOVERY_DB_FILENAME = "discord_message_recovery.db" -_DISCORD_RECOVERY_RETENTION_DAYS = 30 + _DISCORD_COMMAND_SYNC_MUTATION_INTERVAL_SECONDS = 4.5 _DISCORD_COMMAND_SYNC_MAX_RATE_LIMIT_SLEEP_SECONDS = 30.0 # Discord enforces a hard cap of 100 global application (slash) commands per @@ -918,7 +916,9 @@ class DiscordAdapter(BasePlatformAdapter): # shutdown from a runtime websocket crash. self._disconnecting = False self._missed_message_backfill_task: Optional[asyncio.Task] = None - self._discord_recovery_db_lock = threading.Lock() + from hermes_constants import get_hermes_home + from plugins.platforms.discord.recovery import DiscordRecoveryStore + self._discord_recovery_store = DiscordRecoveryStore(get_hermes_home()) # Dedup cache: prevents duplicate bot responses when Discord # RESUME replays events after reconnects. self._dedup = MessageDeduplicator() @@ -1167,14 +1167,7 @@ class DiscordAdapter(BasePlatformAdapter): adapter_self._run_post_connect_initialization() ) if adapter_self._missed_message_backfill_enabled(): - if ( - adapter_self._missed_message_backfill_task - and not adapter_self._missed_message_backfill_task.done() - ): - adapter_self._missed_message_backfill_task.cancel() - adapter_self._missed_message_backfill_task = asyncio.create_task( - adapter_self._run_missed_message_backfill() - ) + adapter_self._ensure_missed_message_backfill_task() @self._client.event async def on_message(message: DiscordMessage): @@ -1253,33 +1246,36 @@ class DiscordAdapter(BasePlatformAdapter): self._release_platform_lock() return False - async def _dispatch_discord_message(self, message: DiscordMessage) -> bool: - """Apply Discord ingress policy and dispatch one live or recovered event.""" - if not self._ready_event.is_set(): - try: - await asyncio.wait_for(self._ready_event.wait(), timeout=30.0) - except asyncio.TimeoutError: - pass - - if self._dedup.is_duplicate(str(message.id)): - return False + def _discord_message_admission( + self, + message: Any, + *, + claim: bool, + ) -> tuple[bool, bool]: + """Return ``(admitted, role_authorized)`` for one Discord event.""" + message_id = str(getattr(message, "id", "")) + if claim: + if self._dedup.is_duplicate(message_id): + return False, False + elif self._dedup.contains(message_id): + return False, False if message.author == self._client.user: - return False + return False, False if message.type not in {discord.MessageType.default, discord.MessageType.reply}: - return False + return False, False role_authorized = False if getattr(message.author, "bot", False): allow_bots = os.getenv("DISCORD_ALLOW_BOTS", "none").lower().strip() if allow_bots == "none": - return False + return False, False if allow_bots == "mentions" and not self._self_is_explicitly_mentioned(message): - return False + return False, False if ( self._discord_bots_require_inline_mention() and not self._self_is_raw_mentioned(message) ): - return False + return False, False else: msg_guild = getattr(message, "guild", None) is_dm = isinstance(message.channel, discord.DMChannel) or msg_guild is None @@ -1297,7 +1293,7 @@ class DiscordAdapter(BasePlatformAdapter): channel_ids=msg_channel_ids, ): self._warn_if_fail_closed_default() - return False + return False, False role_authorized = bool(getattr(self, "_allowed_role_ids", set())) raw_self_mention = self._self_is_explicitly_mentioned(message) @@ -1309,7 +1305,7 @@ class DiscordAdapter(BasePlatformAdapter): for mentioned in message.mentions ) if other_bots_mentioned and not raw_self_mention: - return False + return False, False ignore_no_mention = os.getenv( "DISCORD_IGNORE_NO_MENTION", "true" ).lower() in {"true", "1", "yes"} @@ -1320,10 +1316,25 @@ class DiscordAdapter(BasePlatformAdapter): free_channels = self._discord_free_response_channels() channel_keys = self._discord_channel_keys(message, parent_id) if "*" not in free_channels and not (channel_keys & free_channels): - return False + return False, False - await self._handle_message(message, role_authorized=role_authorized) - return True + return True, role_authorized + + async def _dispatch_discord_message(self, message: Any) -> bool: + """Apply Discord ingress policy and dispatch one live event.""" + if not self._ready_event.is_set(): + try: + await asyncio.wait_for(self._ready_event.wait(), timeout=30.0) + except asyncio.TimeoutError: + pass + admitted, role_authorized = self._discord_message_admission( + message, claim=True, + ) + if not admitted: + return False + return await self._handle_message( + message, role_authorized=role_authorized, + ) async def _cancel_bot_task(self) -> None: """Cancel and await the background client.start() task, if running.""" @@ -1899,6 +1910,12 @@ class DiscordAdapter(BasePlatformAdapter): def _missed_message_backfill_enabled(self) -> bool: """Whether to reconcile Discord messages missed while the gateway was down.""" + configured = self.config.extra.get("missed_message_backfill") + if isinstance(configured, dict) and "enabled" in configured: + value = configured["enabled"] + if isinstance(value, str): + return value.strip().lower() in ("true", "1", "yes", "on") + return bool(value) raw = os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL", "false") return str(raw).strip().lower() in ("true", "1", "yes", "on") @@ -1907,9 +1924,17 @@ class DiscordAdapter(BasePlatformAdapter): Defaults to the union of allowed and free-response channels so both mention-gated requests and mention-free work can be recovered. - Operators can set ``DISCORD_MISSED_MESSAGE_BACKFILL_CHANNELS="*"`` - to scan every reachable text channel, but the safe default is scoped. + Operators can set ``channels: "*"`` to scan every reachable text + channel, but the safe default is scoped. """ + configured = self.config.extra.get("missed_message_backfill") + if isinstance(configured, dict) and "channels" in configured: + raw = configured.get("channels") + if isinstance(raw, list): + return {str(item).strip() for item in raw if str(item).strip()} + raw = str(raw or "") + if raw.strip(): + return {item.strip() for item in raw.split(",") if item.strip()} raw = os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL_CHANNELS", "") if not raw.strip(): allowed = { @@ -1921,7 +1946,12 @@ class DiscordAdapter(BasePlatformAdapter): return {item.strip() for item in raw.split(",") if item.strip()} def _missed_message_backfill_window_seconds(self) -> float: - raw = os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL_WINDOW_SECONDS", "21600") + configured = self.config.extra.get("missed_message_backfill") + raw = ( + configured.get("window_seconds", 21600) + if isinstance(configured, dict) + else os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL_WINDOW_SECONDS", "21600") + ) try: value = float(raw) except (TypeError, ValueError): @@ -1929,7 +1959,12 @@ class DiscordAdapter(BasePlatformAdapter): return max(60.0, value) def _missed_message_backfill_limit(self) -> int: - raw = os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL_LIMIT", "100") + configured = self.config.extra.get("missed_message_backfill") + raw = ( + configured.get("limit", 100) + if isinstance(configured, dict) + else os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL_LIMIT", "100") + ) try: value = int(raw) except (TypeError, ValueError): @@ -1937,13 +1972,34 @@ class DiscordAdapter(BasePlatformAdapter): return max(1, min(value, 500)) def _missed_message_backfill_max_dispatches(self) -> int: - raw = os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL_MAX_DISPATCHES", "10") + configured = self.config.extra.get("missed_message_backfill") + raw = ( + configured.get("max_dispatches", 10) + if isinstance(configured, dict) + else os.getenv("DISCORD_MISSED_MESSAGE_BACKFILL_MAX_DISPATCHES", "10") + ) try: value = int(raw) except (TypeError, ValueError): value = 10 return max(1, min(value, 100)) + def _ensure_missed_message_backfill_task(self) -> asyncio.Task: + """Return the active recovery task, or start one when none is running.""" + task = self._missed_message_backfill_task + if task is not None and not task.done(): + return task + task = asyncio.create_task(self._run_missed_message_backfill()) + self._missed_message_backfill_task = task + runner = getattr(self, "gateway_runner", None) + if runner is not None and getattr(runner, "_startup_restore_in_progress", False): + tasks = getattr(runner, "_startup_restore_tasks", None) + if tasks is None: + tasks = [] + runner._startup_restore_tasks = tasks + tasks.append(task) + return task + async def _run_missed_message_backfill(self) -> None: """Find and enqueue recent Discord messages missed while the bot was down. @@ -1990,7 +2046,12 @@ class DiscordAdapter(BasePlatformAdapter): try: if await self._dispatch_recovered_message(message): dispatched += 1 + except asyncio.CancelledError: + self._dedup.discard(message_id) + self._record_recovery_attempt(message, status="cancelled") + raise except Exception as exc: + self._dedup.discard(message_id) self._record_recovery_attempt(message, status="failed", error=str(exc)) raise if dispatched >= max_dispatches: @@ -2012,7 +2073,33 @@ class DiscordAdapter(BasePlatformAdapter): async def _dispatch_recovered_message(self, message: Any) -> bool: """Run one recovered message through the live Discord ingress gates.""" - return await self._dispatch_discord_message(message) + if not isinstance(message.channel, discord.DMChannel): + parent_id = self._get_parent_channel_id(message.channel) + channel_keys = self._discord_channel_keys(message, parent_id) + free_channels = self._discord_free_response_channels() + in_bot_thread = ( + isinstance(message.channel, discord.Thread) + and str(message.channel.id) in self._threads + and not self._discord_thread_require_mention() + ) + if ( + self._discord_require_mention() + and "*" not in free_channels + and not (channel_keys & free_channels) + and not in_bot_thread + and not self._self_is_explicitly_mentioned(message) + ): + return False + admitted, role_authorized = self._discord_message_admission( + message, claim=False, + ) + if not admitted: + return False + return await self._handle_message( + message, + role_authorized=role_authorized, + recovered=True, + ) async def _iter_missed_message_backfill_candidates(self, channel_ids: set[str]): if not self._client: @@ -2181,74 +2268,17 @@ class DiscordAdapter(BasePlatformAdapter): return False def _discord_recovery_db_path(self) -> _Path: - from hermes_constants import get_hermes_home - - directory = get_hermes_home() / "gateway" - directory.mkdir(parents=True, exist_ok=True) - return directory / _DISCORD_RECOVERY_DB_FILENAME + return self._discord_recovery_store.path() def _with_discord_recovery_db(self, fn, default=None): - try: - with self._discord_recovery_db_lock: - conn = sqlite3.connect(self._discord_recovery_db_path()) - try: - conn.execute("PRAGMA journal_mode=WAL") - conn.execute(""" - CREATE TABLE IF NOT EXISTS discord_messages ( - message_id TEXT PRIMARY KEY, - channel_id TEXT, - thread_id TEXT, - parent_channel_id TEXT, - author_id TEXT, - created_at TEXT, - status TEXT NOT NULL, - replied INTEGER NOT NULL DEFAULT 0, - emoji_ack INTEGER NOT NULL DEFAULT 0, - outage_response INTEGER NOT NULL DEFAULT 0, - response_message_id TEXT, - action_id TEXT, - attempts INTEGER NOT NULL DEFAULT 0, - last_attempt_at TEXT, - last_error TEXT, - updated_at TEXT NOT NULL - ) - """) - conn.execute(""" - CREATE TABLE IF NOT EXISTS discord_recovery_scans ( - scan_id TEXT PRIMARY KEY, - started_at TEXT NOT NULL, - completed_at TEXT, - status TEXT NOT NULL, - channels TEXT NOT NULL, - window_seconds REAL NOT NULL, - limit_count INTEGER NOT NULL, - scanned INTEGER NOT NULL DEFAULT 0, - missed INTEGER NOT NULL DEFAULT 0, - dispatched INTEGER NOT NULL DEFAULT 0, - error TEXT - ) - """) - cutoff = ( - dt.datetime.now(dt.timezone.utc) - - dt.timedelta(days=_DISCORD_RECOVERY_RETENTION_DAYS) - ).isoformat() - conn.execute( - "DELETE FROM discord_messages WHERE updated_at < ?", - (cutoff,), - ) - conn.execute( - "DELETE FROM discord_recovery_scans " - "WHERE COALESCE(completed_at, started_at) < ?", - (cutoff,), - ) - result = fn(conn) - conn.commit() - return result - finally: - conn.close() - except Exception as exc: - logger.debug("[%s] Discord recovery DB operation failed: %s", self.name, exc, exc_info=True) - return default + return self._discord_recovery_store.call(fn, default) + + async def _with_discord_recovery_db_async(self, fn, default=None): + return await asyncio.to_thread( + self._discord_recovery_store.call, + fn, + default, + ) @staticmethod def _utc_now_iso() -> str: @@ -2345,32 +2375,52 @@ class DiscordAdapter(BasePlatformAdapter): def _op(conn): conn.execute( - "UPDATE discord_messages SET status=?, updated_at=? WHERE message_id=?", + "UPDATE discord_messages " + "SET status=CASE WHEN status='responded' THEN status ELSE ? END, " + "updated_at=? WHERE message_id=?", (status, now, message_id), ) self._with_discord_recovery_db(_op) - def _record_discord_response(self, *, reply_to: Optional[str], result: SendResult, content: str) -> None: + def _record_discord_response( + self, + *, + reply_to: Optional[str], + result: SendResult, + content: str, + final: bool, + ) -> None: if not self._missed_message_backfill_enabled() or not reply_to: return now = self._utc_now_iso() - outage = self._is_down_notice_content(content) - status = "missed" if outage else ("responded" if result.success else "failed") + completed = bool(final and result.success) + status = "responded" if completed else "failed" def _op(conn): conn.execute( """ INSERT INTO discord_messages (message_id, status, replied, outage_response, response_message_id, updated_at) - VALUES (?, ?, ?, ?, ?, ?) + VALUES (?, ?, ?, 0, ?, ?) ON CONFLICT(message_id) DO UPDATE SET - status=?, - replied=CASE WHEN ? THEN 1 ELSE replied END, - outage_response=CASE WHEN ? THEN 1 ELSE outage_response END, + status=CASE WHEN ? THEN 'responded' ELSE discord_messages.status END, + replied=CASE WHEN ? THEN 1 ELSE discord_messages.replied END, + outage_response=CASE WHEN ? THEN 0 ELSE discord_messages.outage_response END, response_message_id=COALESCE(?, response_message_id), updated_at=? """, - (reply_to, status, 1 if result.success and not outage else 0, 1 if outage else 0, result.message_id, now, status, 1 if result.success and not outage else 0, 1 if outage else 0, result.message_id, now), + ( + reply_to, + status, + 1 if completed else 0, + result.message_id, + now, + 1 if completed else 0, + 1 if completed else 0, + 1 if completed else 0, + result.message_id, + now, + ), ) self._with_discord_recovery_db(_op) @@ -2639,11 +2689,19 @@ class DiscordAdapter(BasePlatformAdapter): acked = False if self._reactions_enabled() and hasattr(message, "add_reaction"): acked = await self._add_reaction(message, "👀") - self._record_discord_processing_start(event, emoji_ack=acked) + await asyncio.to_thread( + self._record_discord_processing_start, + event, + emoji_ack=acked, + ) async def on_processing_complete(self, event: MessageEvent, outcome: ProcessingOutcome) -> None: """Swap the in-progress reaction for final reaction and durable state.""" - self._record_discord_processing_complete(event, outcome) + await asyncio.to_thread( + self._record_discord_processing_complete, + event, + outcome, + ) if not self._reactions_enabled(): return message = event.raw_message @@ -2678,6 +2736,7 @@ class DiscordAdapter(BasePlatformAdapter): if metadata and metadata.get("thread_id"): thread_id = metadata["thread_id"] nonconversational = _metadata_marks_nonconversational(metadata) + final_delivery = bool(metadata and metadata.get("notify")) if thread_id: # Fetch the thread directly — threads are addressed by their own ID. @@ -2701,6 +2760,7 @@ class DiscordAdapter(BasePlatformAdapter): reply_to=reply_to, result=result, content=content, + final=final_delivery, ) return result @@ -2771,13 +2831,23 @@ class DiscordAdapter(BasePlatformAdapter): message_id=message_ids[0] if message_ids else None, raw_response={"message_ids": message_ids} ) - self._record_discord_response(reply_to=reply_to, result=result, content=content) + self._record_discord_response( + reply_to=reply_to, + result=result, + content=content, + final=final_delivery, + ) return result except Exception as e: # pragma: no cover - defensive logging logger.error("[%s] Failed to send Discord message: %s", self.name, e, exc_info=True) result = SendResult(success=False, error=str(e)) - self._record_discord_response(reply_to=reply_to, result=result, content=content) + self._record_discord_response( + reply_to=reply_to, + result=result, + content=content, + final=bool(metadata and metadata.get("notify")), + ) return result async def _send_to_forum(self, forum_channel: Any, content: str) -> SendResult: @@ -6870,8 +6940,14 @@ class DiscordAdapter(BasePlatformAdapter): raise Exception(f"HTTP {resp.status}") return await resp.read() - async def _handle_message(self, message: DiscordMessage, role_authorized: bool = False) -> None: - """Handle incoming Discord messages.""" + async def _handle_message( + self, + message: DiscordMessage, + role_authorized: bool = False, + *, + recovered: bool = False, + ) -> bool: + """Handle one Discord message and report whether it reached dispatch.""" # In server channels (not DMs), require the bot to be @mentioned # UNLESS the channel is in the free-response list or the message is # in a thread where the bot has already participated. @@ -6927,14 +7003,14 @@ class DiscordAdapter(BasePlatformAdapter): allowed_channels = {ch.strip() for ch in allowed_channels_raw.split(",") if ch.strip()} if "*" not in allowed_channels and not (channel_keys & allowed_channels): logger.debug("[%s] Ignoring message in non-allowed channel: %s", self.name, channel_keys) - return + return False # Check ignored channels - never respond even when mentioned ignored_channels_raw = os.getenv("DISCORD_IGNORED_CHANNELS", "") ignored_channels = {ch.strip() for ch in ignored_channels_raw.split(",") if ch.strip()} if "*" in ignored_channels or (channel_keys & ignored_channels): logger.debug("[%s] Ignoring message in ignored channel: %s", self.name, channel_keys) - return + return False free_channels = self._discord_free_response_channels() @@ -6963,7 +7039,7 @@ class DiscordAdapter(BasePlatformAdapter): if require_mention and not is_free_channel and not in_bot_thread: if not self._self_is_explicitly_mentioned(message) and not mention_prefix: - return + return False # Auto-thread: when enabled, automatically create a thread for every # @mention in a text channel so each conversation is isolated (like Slack). # Messages already inside threads or DMs are unaffected. @@ -7012,7 +7088,7 @@ class DiscordAdapter(BasePlatformAdapter): self.name, notify_error, ) - return + return False referenced_attachments = [] reference = getattr(message, "reference", None) @@ -7309,7 +7385,7 @@ class DiscordAdapter(BasePlatformAdapter): getattr(message.author, "display_name", getattr(message.author, "name", "unknown")), getattr(message.channel, "id", "unknown"), ) - return + return False event_text = "(The user sent a message with no text content)" _chan = message.channel @@ -7346,12 +7422,18 @@ class DiscordAdapter(BasePlatformAdapter): if thread_id: self._threads.mark(thread_id) - # Only batch plain text messages — commands, media, etc. dispatch - # immediately since they won't be split by the Discord client. - if msg_type == MessageType.TEXT and self._text_batch_delay_seconds > 0: + # Only live plain text messages use split-message batching. Recovery + # candidates are already complete historical messages; coalescing them + # would lose constituent IDs and make later restarts replay them. + if ( + not recovered + and msg_type == MessageType.TEXT + and self._text_batch_delay_seconds > 0 + ): self._enqueue_text_event(event) else: await self.handle_message(event) + return True # ------------------------------------------------------------------ # Text message aggregation (handles Discord client-side splits) @@ -9147,22 +9229,10 @@ def _apply_yaml_config(yaml_cfg: dict, discord_cfg: dict) -> dict | None: os.environ["DISCORD_AUTO_THREAD"] = str(discord_cfg["auto_thread"]).lower() if "reactions" in discord_cfg and not os.getenv("DISCORD_REACTIONS"): os.environ["DISCORD_REACTIONS"] = str(discord_cfg["reactions"]).lower() + seeded_extra = {} backfill_cfg = discord_cfg.get("missed_message_backfill") if isinstance(backfill_cfg, dict): - _backfill_env = { - "enabled": "DISCORD_MISSED_MESSAGE_BACKFILL", - "channels": "DISCORD_MISSED_MESSAGE_BACKFILL_CHANNELS", - "window_seconds": "DISCORD_MISSED_MESSAGE_BACKFILL_WINDOW_SECONDS", - "limit": "DISCORD_MISSED_MESSAGE_BACKFILL_LIMIT", - "max_dispatches": "DISCORD_MISSED_MESSAGE_BACKFILL_MAX_DISPATCHES", - } - for _key, _env in _backfill_env.items(): - if _key not in backfill_cfg or os.getenv(_env): - continue - _value = backfill_cfg[_key] - if isinstance(_value, list): - _value = ",".join(str(v) for v in _value) - os.environ[_env] = str(_value).lower() if isinstance(_value, bool) else str(_value) + seeded_extra["missed_message_backfill"] = dict(backfill_cfg) # ignored_channels: channels where bot never responds (even when mentioned) ic = discord_cfg.get("ignored_channels") if ic is not None and not os.getenv("DISCORD_IGNORED_CHANNELS"): @@ -9239,16 +9309,15 @@ def _apply_yaml_config(yaml_cfg: dict, discord_cfg: dict) -> dict | None: ("websocket_heartbeat_ack_max_age_seconds", None, None), ("websocket_max_latency_seconds", None, None), ) - seeded = {} for primary_key, legacy_key, env_key in _websocket_liveness_keys: value = _websocket_liveness_cfg.get(primary_key) if value is None and legacy_key: value = _websocket_liveness_cfg.get(legacy_key) if value is not None: - seeded[primary_key] = value + seeded_extra[primary_key] = value if env_key and not os.getenv(env_key): os.environ[env_key] = str(value) - return seeded or None + return seeded_extra or None def _is_connected(config) -> bool: diff --git a/plugins/platforms/discord/recovery.py b/plugins/platforms/discord/recovery.py new file mode 100644 index 000000000..cac4072b6 --- /dev/null +++ b/plugins/platforms/discord/recovery.py @@ -0,0 +1,99 @@ +"""Durable state for Discord reconnect message recovery.""" + +from __future__ import annotations + +import datetime as dt +import logging +import os +import sqlite3 +import threading +from contextlib import suppress +from pathlib import Path +from typing import Any, Callable + +from hermes_constants import get_hermes_home + +logger = logging.getLogger(__name__) + +_DB_FILENAME = "discord_message_recovery.db" +_RETENTION_DAYS = 30 + + +class DiscordRecoveryStore: + """Small profile-scoped SQLite ledger for completed Discord messages.""" + + def __init__(self, hermes_home: Path | None = None) -> None: + self._lock = threading.Lock() + self._initialized = False + self._hermes_home = Path(hermes_home or get_hermes_home()) + + def path(self) -> Path: + directory = self._hermes_home / "gateway" + directory.mkdir(parents=True, exist_ok=True) + return directory / _DB_FILENAME + + def call(self, fn: Callable[[sqlite3.Connection], Any], default: Any = None) -> Any: + try: + with self._lock: + path = self.path() + conn = sqlite3.connect(path, timeout=0.1) + try: + if not self._initialized: + self._initialize(conn) + self._initialized = True + with suppress(OSError): + os.chmod(path, 0o600) + result = fn(conn) + conn.commit() + return result + finally: + conn.close() + except Exception as exc: + logger.warning("Discord recovery ledger unavailable: %s", exc) + return default + + def _initialize(self, conn: sqlite3.Connection) -> None: + conn.execute("PRAGMA journal_mode=WAL") + conn.execute(""" + CREATE TABLE IF NOT EXISTS discord_messages ( + message_id TEXT PRIMARY KEY, + channel_id TEXT, + thread_id TEXT, + parent_channel_id TEXT, + author_id TEXT, + created_at TEXT, + status TEXT NOT NULL, + replied INTEGER NOT NULL DEFAULT 0, + emoji_ack INTEGER NOT NULL DEFAULT 0, + outage_response INTEGER NOT NULL DEFAULT 0, + response_message_id TEXT, + attempts INTEGER NOT NULL DEFAULT 0, + last_attempt_at TEXT, + last_error TEXT, + updated_at TEXT NOT NULL + ) + """) + conn.execute(""" + CREATE TABLE IF NOT EXISTS discord_recovery_scans ( + scan_id TEXT PRIMARY KEY, + started_at TEXT NOT NULL, + completed_at TEXT, + status TEXT NOT NULL, + channels TEXT NOT NULL, + window_seconds REAL NOT NULL, + limit_count INTEGER NOT NULL, + scanned INTEGER NOT NULL DEFAULT 0, + missed INTEGER NOT NULL DEFAULT 0, + dispatched INTEGER NOT NULL DEFAULT 0, + error TEXT + ) + """) + cutoff = ( + dt.datetime.now(dt.timezone.utc) - dt.timedelta(days=_RETENTION_DAYS) + ).isoformat() + conn.execute("DELETE FROM discord_messages WHERE updated_at < ?", (cutoff,)) + conn.execute( + "DELETE FROM discord_recovery_scans " + "WHERE COALESCE(completed_at, started_at) < ?", + (cutoff,), + ) diff --git a/tests/gateway/test_discord_missed_message_backfill.py b/tests/gateway/test_discord_missed_message_backfill.py index 93c635481..bd4d2ff67 100644 --- a/tests/gateway/test_discord_missed_message_backfill.py +++ b/tests/gateway/test_discord_missed_message_backfill.py @@ -1,5 +1,6 @@ """Tests for Discord missed-message startup backfill.""" +import asyncio import datetime as dt import os import sys @@ -9,7 +10,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from gateway.config import PlatformConfig +from gateway.config import Platform, PlatformConfig from gateway.platforms.base import MessageEvent, MessageType, ProcessingOutcome @@ -48,7 +49,10 @@ def _ensure_discord_mock(): _ensure_discord_mock() import discord # noqa: E402 -from plugins.platforms.discord.adapter import DiscordAdapter # noqa: E402 +from plugins.platforms.discord.adapter import ( # noqa: E402 + DiscordAdapter, + _apply_yaml_config, +) class FakeReaction: @@ -67,7 +71,7 @@ class FakeChannel: self.id = channel_id self.parent_id = parent_id self.name = "wiki-inbox" - self.guild = SimpleNamespace(name="emo") + self.guild = SimpleNamespace(id=777, name="emo") self.topic = None self._history_messages = list(history_messages or []) @@ -87,7 +91,7 @@ def adapter(monkeypatch, tmp_path): bot_user = SimpleNamespace(id=999, bot=True, display_name="Hermes", name="hermes") adapter._client = SimpleNamespace(user=bot_user, get_channel=lambda _id: None) adapter._ready_event.set() - adapter._handle_message = AsyncMock() + adapter._handle_message = AsyncMock(return_value=True) monkeypatch.setenv("DISCORD_MISSED_MESSAGE_BACKFILL", "true") monkeypatch.setenv("DISCORD_ALLOW_ALL_USERS", "true") return adapter @@ -194,7 +198,12 @@ async def test_generic_unavailable_response_counts_as_completed(adapter): @pytest.mark.asyncio async def test_run_backfill_dispatches_unaddressed_messages(adapter, monkeypatch): - message = make_message(message_id=1) + bot_user = adapter._client.user + message = make_message( + message_id=1, + content=f"<@{bot_user.id}> please ingest", + mentions=[bot_user], + ) async def fake_candidates(_channels): yield message @@ -208,7 +217,11 @@ async def test_run_backfill_dispatches_unaddressed_messages(adapter, monkeypatch await adapter._run_missed_message_backfill() - adapter._handle_message.assert_awaited_once_with(message, role_authorized=False) + adapter._handle_message.assert_awaited_once_with( + message, + role_authorized=False, + recovered=True, + ) @pytest.mark.asyncio @@ -235,6 +248,75 @@ async def test_run_backfill_counts_only_messages_that_reach_dispatch(adapter, mo assert dispatch.await_count == 2 +@pytest.mark.asyncio +async def test_recovery_releases_dedup_claim_when_dispatch_is_cancelled(adapter, monkeypatch): + message = make_message(message_id=97) + started = asyncio.Event() + + async def cancelled_dispatch(_message): + adapter._dedup.is_duplicate(str(message.id)) + started.set() + await asyncio.Event().wait() + + monkeypatch.setattr(adapter, "_dispatch_recovered_message", cancelled_dispatch) + monkeypatch.setattr(adapter, "_should_backfill_discord_message", AsyncMock(return_value=True)) + monkeypatch.setattr(adapter, "_missed_message_backfill_channels", lambda: {"123"}) + + async def candidates(_channels): + yield message + + monkeypatch.setattr(adapter, "_iter_missed_message_backfill_candidates", candidates) + task = asyncio.create_task(adapter._run_missed_message_backfill()) + await started.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert adapter._dedup.contains(str(message.id)) is False + + +@pytest.mark.asyncio +async def test_repeated_ready_coalesces_instead_of_cancelling_active_recovery(adapter): + started = asyncio.Event() + release = asyncio.Event() + + async def slow_recovery(): + started.set() + await release.wait() + + first = asyncio.create_task(slow_recovery()) + adapter._missed_message_backfill_task = first + await started.wait() + + second = adapter._ensure_missed_message_backfill_task() + + assert second is first + assert first.cancelled() is False + release.set() + await first + + +@pytest.mark.asyncio +async def test_recovery_task_joins_gateway_startup_restore(adapter, monkeypatch): + release = asyncio.Event() + + async def recovery(): + await release.wait() + + runner = SimpleNamespace( + _startup_restore_in_progress=True, + _startup_restore_tasks=[], + ) + adapter.gateway_runner = runner + monkeypatch.setattr(adapter, "_run_missed_message_backfill", recovery) + + task = adapter._ensure_missed_message_backfill_task() + + assert runner._startup_restore_tasks == [task] + release.set() + await task + + @pytest.mark.asyncio async def test_recovered_mention_reuses_live_auth_and_mention_gates(adapter, monkeypatch): bot_user = adapter._client.user @@ -259,7 +341,43 @@ async def test_recovered_mention_reuses_live_auth_and_mention_gates(adapter, mon assert await adapter._dispatch_recovered_message(denied) is False assert await adapter._dispatch_recovered_message(allowed) is True - adapter._handle_message.assert_awaited_once_with(allowed, role_authorized=False) + adapter._handle_message.assert_awaited_once_with( + allowed, + role_authorized=False, + recovered=True, + ) + + +@pytest.mark.asyncio +async def test_recovery_does_not_treat_unmentioned_message_as_dispatched(adapter, monkeypatch): + monkeypatch.setenv("DISCORD_REQUIRE_MENTION", "true") + monkeypatch.setenv("DISCORD_AUTO_THREAD", "false") + adapter.config.extra["free_response_channels"] = "" + adapter.handle_message = AsyncMock() + message = make_message(message_id=95, content="not addressed") + + assert await adapter._dispatch_recovered_message(message) is False + adapter.handle_message.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_recovered_messages_bypass_live_text_debounce(adapter, monkeypatch): + bot_user = adapter._client.user + message = make_message( + message_id=96, + content=f"<@{bot_user.id}> recover", + mentions=[bot_user], + ) + adapter._text_batch_delay_seconds = 0.6 + adapter._handle_message = DiscordAdapter._handle_message.__get__( + adapter, DiscordAdapter + ) + adapter.handle_message = AsyncMock() + monkeypatch.setenv("DISCORD_AUTO_THREAD", "false") + + assert await adapter._dispatch_recovered_message(message) is True + adapter.handle_message.assert_awaited_once() + assert adapter._pending_text_batches == {} def test_missed_message_backfill_config_bridge(monkeypatch, tmp_path): @@ -288,13 +406,18 @@ def test_missed_message_backfill_config_bridge(monkeypatch, tmp_path): " max_dispatches: 3\n" ) - load_gateway_config() + config = load_gateway_config() + backfill = config.platforms[Platform.DISCORD].extra[ + "missed_message_backfill" + ] - assert os.environ["DISCORD_MISSED_MESSAGE_BACKFILL"] == "true" - assert os.environ["DISCORD_MISSED_MESSAGE_BACKFILL_CHANNELS"] == "1501971993405292796" - assert os.environ["DISCORD_MISSED_MESSAGE_BACKFILL_WINDOW_SECONDS"] == "3600" - assert os.environ["DISCORD_MISSED_MESSAGE_BACKFILL_LIMIT"] == "25" - assert os.environ["DISCORD_MISSED_MESSAGE_BACKFILL_MAX_DISPATCHES"] == "3" + assert backfill == { + "enabled": True, + "channels": ["1501971993405292796"], + "window_seconds": 3600, + "limit": 25, + "max_dispatches": 3, + } def test_default_config_exposes_missed_message_backfill_settings(): @@ -309,6 +432,59 @@ def test_default_config_exposes_missed_message_backfill_settings(): } +def test_missed_message_backfill_config_stays_per_adapter(): + first_extra = _apply_yaml_config( + {}, + { + "missed_message_backfill": { + "enabled": True, + "channels": ["111"], + "window_seconds": 60, + "limit": 5, + "max_dispatches": 2, + } + }, + ) + second_extra = _apply_yaml_config( + {}, + { + "missed_message_backfill": { + "enabled": False, + "channels": ["222"], + "window_seconds": 120, + "limit": 6, + "max_dispatches": 3, + } + }, + ) + + first = DiscordAdapter(PlatformConfig(enabled=True, token="one", extra=first_extra or {})) + second = DiscordAdapter(PlatformConfig(enabled=True, token="two", extra=second_extra or {})) + + assert first._missed_message_backfill_enabled() is True + assert first._missed_message_backfill_channels() == {"111"} + assert first._missed_message_backfill_window_seconds() == 60 + assert first._missed_message_backfill_limit() == 5 + assert first._missed_message_backfill_max_dispatches() == 2 + assert second._missed_message_backfill_enabled() is False + assert second._missed_message_backfill_channels() == {"222"} + assert second._missed_message_backfill_window_seconds() == 120 + assert second._missed_message_backfill_limit() == 6 + assert second._missed_message_backfill_max_dispatches() == 3 + + +def test_recovery_store_pins_profile_home_at_adapter_construction(monkeypatch, tmp_path): + first_home = tmp_path / "first" + second_home = tmp_path / "second" + monkeypatch.setenv("HERMES_HOME", str(first_home)) + adapter = DiscordAdapter(PlatformConfig(enabled=True, token="one")) + monkeypatch.setenv("HERMES_HOME", str(second_home)) + + assert adapter._discord_recovery_db_path() == ( + first_home / "gateway" / "discord_message_recovery.db" + ) + + def test_default_recovery_scope_includes_allowed_and_free_response_channels(adapter, monkeypatch): monkeypatch.delenv("DISCORD_MISSED_MESSAGE_BACKFILL_CHANNELS", raising=False) monkeypatch.setenv("DISCORD_ALLOWED_CHANNELS", "100,200") @@ -325,6 +501,7 @@ async def test_persistent_responded_record_suppresses_backfill(adapter): reply_to="77", result=SimpleNamespace(success=True, message_id="9001"), content="Done — captured it.", + final=True, ) assert await adapter._should_backfill_discord_message(message) is False @@ -333,8 +510,9 @@ async def test_persistent_responded_record_suppresses_backfill(adapter): def test_down_notice_response_does_not_mark_message_complete(adapter): adapter._record_discord_response( reply_to="88", - result=SimpleNamespace(success=True, message_id="9002"), + result=SimpleNamespace(success=False, message_id="9002"), content="The agent is down right now.", + final=True, ) assert adapter._discord_message_is_persistently_complete("88") is False @@ -357,6 +535,7 @@ def test_recovery_ledger_prunes_expired_rows(adapter): ) adapter._with_discord_recovery_db(insert_old_rows) + adapter._discord_recovery_store._initialized = False adapter._with_discord_recovery_db(lambda _conn: None) def count_old(conn): @@ -385,6 +564,83 @@ def test_empty_successful_turn_is_not_persistently_complete(adapter): assert adapter._discord_message_is_persistently_complete("89") is False +def test_final_delivery_remains_complete_after_processing_hook(adapter): + message = make_message(message_id=91) + event = MessageEvent( + text=message.content, + message_type=MessageType.TEXT, + raw_message=message, + message_id=str(message.id), + ) + + adapter._record_discord_processing_start(event, emoji_ack=False) + adapter._record_discord_response( + reply_to="91", + result=SimpleNamespace(success=True, message_id="9004"), + content="Done", + final=True, + ) + adapter._record_discord_processing_complete(event, ProcessingOutcome.SUCCESS) + + assert adapter._discord_message_is_persistently_complete("91") is True + + +def test_preview_delivery_does_not_mark_message_complete(adapter): + adapter._record_discord_response( + reply_to="92", + result=SimpleNamespace(success=True, message_id="9005"), + content="partial", + final=False, + ) + + assert adapter._discord_message_is_persistently_complete("92") is False + + +def test_successful_final_delivery_clears_prior_outage_state(adapter): + adapter._record_discord_response( + reply_to="93", + result=SimpleNamespace(success=False, message_id="9006"), + content="Hermes is offline", + final=True, + ) + assert adapter._discord_message_is_persistently_complete("93") is False + + adapter._record_discord_response( + reply_to="93", + result=SimpleNamespace(success=True, message_id="9007"), + content="Recovered successfully", + final=True, + ) + + assert adapter._discord_message_is_persistently_complete("93") is True + + +@pytest.mark.asyncio +async def test_send_uses_notify_metadata_as_final_delivery_signal(adapter): + channel = FakeChannel(channel_id=123) + channel.send = AsyncMock(return_value=SimpleNamespace(id=9008)) + channel.fetch_message = AsyncMock() + adapter._client.get_channel = lambda _channel_id: channel + + preview = await adapter.send( + "123", + "partial", + reply_to="94", + metadata={"expect_edits": True}, + ) + assert preview.success is True + assert adapter._discord_message_is_persistently_complete("94") is False + + final = await adapter.send( + "123", + "complete", + reply_to="94", + metadata={"notify": True}, + ) + assert final.success is True + assert adapter._discord_message_is_persistently_complete("94") is True + + def test_disabled_recovery_does_not_create_hot_path_ledger(adapter, monkeypatch): monkeypatch.setenv("DISCORD_MISSED_MESSAGE_BACKFILL", "false") message = make_message(message_id=90) @@ -401,6 +657,7 @@ def test_disabled_recovery_does_not_create_hot_path_ledger(adapter, monkeypatch) reply_to="90", result=SimpleNamespace(success=True, message_id="9003"), content="Done", + final=True, ) db_path = adapter._discord_recovery_db_path()