diff --git a/gateway/config.py b/gateway/config.py index bde52eb5596a..530d3d7f3b4a 100644 --- a/gateway/config.py +++ b/gateway/config.py @@ -754,6 +754,38 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if Platform.DISCORD not in config.platforms: config.platforms[Platform.DISCORD] = PlatformConfig() config.platforms[Platform.DISCORD].reply_to_mode = discord_reply_mode + + discord_allow_bots = os.getenv("DISCORD_ALLOW_BOTS", "").lower().strip() + if discord_allow_bots in ("none", "mentions", "all"): + if Platform.DISCORD not in config.platforms: + config.platforms[Platform.DISCORD] = PlatformConfig() + config.platforms[Platform.DISCORD].extra["allow_bots"] = discord_allow_bots + + discord_target_bot_ids = os.getenv("DISCORD_TARGET_BOT_IDS", "").strip() + if discord_target_bot_ids: + if Platform.DISCORD not in config.platforms: + config.platforms[Platform.DISCORD] = PlatformConfig() + config.platforms[Platform.DISCORD].extra["target_bot_ids"] = [ + bot_id.strip() for bot_id in discord_target_bot_ids.split(",") if bot_id.strip() + ] + + discord_auto_mention_target_bots = os.getenv("DISCORD_AUTO_MENTION_TARGET_BOTS", "") + if discord_auto_mention_target_bots: + if Platform.DISCORD not in config.platforms: + config.platforms[Platform.DISCORD] = PlatformConfig() + config.platforms[Platform.DISCORD].extra["auto_mention_target_bots"] = _coerce_bool( + discord_auto_mention_target_bots, + default=False, + ) + + discord_accept_bot_replies = os.getenv("DISCORD_ACCEPT_BOT_REPLIES", "") + if discord_accept_bot_replies: + if Platform.DISCORD not in config.platforms: + config.platforms[Platform.DISCORD] = PlatformConfig() + config.platforms[Platform.DISCORD].extra["accept_bot_replies"] = _coerce_bool( + discord_accept_bot_replies, + default=True, + ) # WhatsApp (typically uses different auth mechanism) whatsapp_enabled = os.getenv("WHATSAPP_ENABLED", "").lower() in ("true", "1", "yes") diff --git a/gateway/platforms/discord.py b/gateway/platforms/discord.py index dcf05a16257b..865601a22eeb 100644 --- a/gateway/platforms/discord.py +++ b/gateway/platforms/discord.py @@ -76,6 +76,50 @@ def _clean_discord_id(entry: str) -> str: return entry.strip() +def _coerce_bool(value: Any, default: bool = False) -> bool: + """Coerce bool-ish platform config values.""" + if value is None: + return default + if isinstance(value, bool): + return value + if isinstance(value, str): + lowered = value.strip().lower() + if lowered in ("true", "1", "yes", "on"): + return True + if lowered in ("false", "0", "no", "off"): + return False + return default + return bool(value) + + +def _normalize_allow_bots_mode(value: Any) -> str: + """Normalize Discord bot-ingress mode to none|mentions|all.""" + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"none", "mentions", "all"}: + return normalized + return "none" + + +def _normalize_discord_id_list(value: Any) -> tuple[str, ...]: + """Normalize config/env/message metadata into a tuple of Discord IDs.""" + if value is None: + return () + if isinstance(value, str): + items = value.split(",") + elif isinstance(value, (list, tuple, set)): + items = value + else: + items = [value] + + cleaned: list[str] = [] + for item in items: + entry = _clean_discord_id(str(item)) + if entry and entry not in cleaned: + cleaned.append(entry) + return tuple(cleaned) + + def check_discord_requirements() -> bool: """Check if Discord dependencies are available.""" return DISCORD_AVAILABLE @@ -465,6 +509,94 @@ def __init__(self, config: PlatformConfig): # Reply threading mode: "off" (no replies), "first" (reply on first # chunk only, default), "all" (reply-reference on every chunk). self._reply_to_mode: str = getattr(config, 'reply_to_mode', 'first') or 'first' + self._allow_bots_mode: str = _normalize_allow_bots_mode( + self.config.extra.get("allow_bots", os.getenv("DISCORD_ALLOW_BOTS", "none")) + ) + self._accept_bot_replies: bool = _coerce_bool( + self.config.extra.get("accept_bot_replies", os.getenv("DISCORD_ACCEPT_BOT_REPLIES")), + default=True, + ) + self._target_bot_ids: tuple[str, ...] = _normalize_discord_id_list( + self.config.extra.get("target_bot_ids", os.getenv("DISCORD_TARGET_BOT_IDS", "")) + ) + self._auto_mention_target_bots: bool = _coerce_bool( + self.config.extra.get( + "auto_mention_target_bots", + os.getenv("DISCORD_AUTO_MENTION_TARGET_BOTS"), + ), + default=False, + ) + + def _is_reply_to_own_message(self, message: DiscordMessage) -> bool: + """Return True when a Discord message is replying to this bot.""" + if not self._client or not self._client.user: + return False + reference = getattr(message, "reference", None) + if not reference: + return False + resolved = getattr(reference, "resolved", None) or getattr(reference, "cached_message", None) + if resolved is None: + return False + author = getattr(resolved, "author", None) + if author is None: + return False + author_id = getattr(author, "id", None) + self_id = getattr(self._client.user, "id", None) + return author == self._client.user or (author_id is not None and author_id == self_id) + + def _should_accept_bot_message(self, message: DiscordMessage) -> bool: + """Apply bot-ingress policy for messages authored by other bots.""" + if not getattr(message.author, "bot", False): + return True + if self._allow_bots_mode == "all": + return True + + self_user = self._client.user if self._client else None + self_mentioned = bool(self_user and self_user in getattr(message, "mentions", [])) + if self_mentioned: + return True + + if self._accept_bot_replies and self._is_reply_to_own_message(message): + return True + + return False + + def _resolve_outbound_bot_mentions( + self, + metadata: Optional[Dict[str, Any]] = None, + ) -> tuple[str, ...]: + """Return bot IDs that should be @mentioned on outbound messages.""" + metadata = metadata or {} + explicit_ids = _normalize_discord_id_list(metadata.get("mention_bot_ids")) + if explicit_ids: + return explicit_ids + if metadata.get("mention_target_bots"): + return self._target_bot_ids + if self._auto_mention_target_bots: + return self._target_bot_ids + return () + + def _prefix_outbound_bot_mentions(self, content: str, bot_ids: tuple[str, ...]) -> str: + """Prefix any missing bot mentions onto the outgoing content.""" + if not bot_ids: + return content + missing = [ + bot_id for bot_id in bot_ids + if f"<@{bot_id}>" not in content and f"<@!{bot_id}>" not in content + ] + if not missing: + return content + prefix = " ".join(f"<@{bot_id}>" for bot_id in missing) + return f"{prefix} {content}".strip() + + def _build_allowed_mentions(self, bot_ids: tuple[str, ...]): + """Allow explicit user mentions while blocking @everyone/@here escalation.""" + if not bot_ids or discord is None or not hasattr(discord, "AllowedMentions"): + return None + try: + return discord.AllowedMentions(users=True, roles=False, everyone=False, replied_user=False) + except Exception: + return None async def connect(self) -> bool: """Connect to Discord and start receiving events.""" @@ -594,18 +726,13 @@ async def on_message(message: DiscordMessage): if not self._is_allowed_user(str(message.author.id)): return - # Bot message filtering (DISCORD_ALLOW_BOTS): + # Bot message filtering (config.extra.allow_bots / DISCORD_ALLOW_BOTS): # "none" — ignore all other bots (default) - # "mentions" — accept bot messages only when they @mention us + # "mentions" — accept bot messages that @mention us, plus + # reply-based follow-ups when accept_bot_replies is enabled # "all" — accept all bot messages - if getattr(message.author, "bot", False): - allow_bots = os.getenv("DISCORD_ALLOW_BOTS", "none").lower().strip() - if allow_bots == "none": - return - elif allow_bots == "mentions": - if not self._client.user or self._client.user not in message.mentions: - return - # "all" falls through to handle_message + if not self._should_accept_bot_message(message): + return # Multi-agent filtering: if the message mentions specific bots # but NOT this bot, the sender is talking to another agent — @@ -818,7 +945,10 @@ async def send( # Format and split message if needed formatted = self.format_message(content) + outbound_bot_mentions = self._resolve_outbound_bot_mentions(metadata) + formatted = self._prefix_outbound_bot_mentions(formatted, outbound_bot_mentions) chunks = self.truncate_message(formatted, self.MAX_MESSAGE_LENGTH) + allowed_mentions = self._build_allowed_mentions(outbound_bot_mentions) message_ids = [] reference = None @@ -836,10 +966,13 @@ async def send( else: # "first" (default) or "off" chunk_reference = reference if i == 0 else None try: - msg = await channel.send( - content=chunk, - reference=chunk_reference, - ) + send_kwargs = { + "content": chunk, + "reference": chunk_reference, + } + if allowed_mentions is not None: + send_kwargs["allowed_mentions"] = allowed_mentions + msg = await channel.send(**send_kwargs) except Exception as e: err_text = str(e) if ( @@ -852,10 +985,13 @@ async def send( self.name, reply_to, ) - msg = await channel.send( - content=chunk, - reference=None, - ) + retry_kwargs = { + "content": chunk, + "reference": None, + } + if allowed_mentions is not None: + retry_kwargs["allowed_mentions"] = allowed_mentions + msg = await channel.send(**retry_kwargs) else: raise message_ids.append(str(msg.id)) diff --git a/tests/gateway/test_discord_bot_filter.py b/tests/gateway/test_discord_bot_filter.py index 09a78ae6308f..e70c72775f16 100644 --- a/tests/gateway/test_discord_bot_filter.py +++ b/tests/gateway/test_discord_bot_filter.py @@ -1,9 +1,12 @@ -"""Tests for Discord bot message filtering (DISCORD_ALLOW_BOTS).""" +"""Tests for Discord bot message filtering and bot-reply intake.""" -import asyncio import os import unittest -from unittest.mock import AsyncMock, MagicMock, patch +from types import SimpleNamespace +from unittest.mock import MagicMock + +from gateway.config import PlatformConfig +from gateway.platforms.discord import DiscordAdapter def _make_author(*, bot: bool = False, is_self: bool = False): @@ -113,5 +116,43 @@ def test_case_insensitive(self): self.assertFalse(self._run_filter(msg, "None")) +class TestDiscordBotReplyIntake(unittest.TestCase): + """Test adapter-level bot message intake for reply-based handoffs.""" + + def _make_adapter(self, *, allow_bots="mentions", accept_bot_replies=True): + config = PlatformConfig( + enabled=True, + token="test-token", + extra={ + "allow_bots": allow_bots, + "accept_bot_replies": accept_bot_replies, + }, + ) + adapter = DiscordAdapter(config) + adapter._client = SimpleNamespace(user=SimpleNamespace(id=99999, bot=True)) + return adapter + + def test_bot_reply_to_self_is_accepted_without_mention(self): + adapter = self._make_adapter() + reply_target = SimpleNamespace(author=adapter._client.user) + msg = _make_message(author=_make_author(bot=True), mentions=[]) + msg.reference = SimpleNamespace(resolved=reply_target) + self.assertTrue(adapter._should_accept_bot_message(msg)) + + def test_bot_reply_to_other_user_is_rejected_without_mention(self): + adapter = self._make_adapter() + reply_target = SimpleNamespace(author=_make_author(bot=False)) + msg = _make_message(author=_make_author(bot=True), mentions=[]) + msg.reference = SimpleNamespace(resolved=reply_target) + self.assertFalse(adapter._should_accept_bot_message(msg)) + + def test_bot_reply_acceptance_can_be_disabled(self): + adapter = self._make_adapter(accept_bot_replies=False) + reply_target = SimpleNamespace(author=adapter._client.user) + msg = _make_message(author=_make_author(bot=True), mentions=[]) + msg.reference = SimpleNamespace(resolved=reply_target) + self.assertFalse(adapter._should_accept_bot_message(msg)) + + if __name__ == "__main__": unittest.main() diff --git a/tests/gateway/test_discord_reply_mode.py b/tests/gateway/test_discord_reply_mode.py index 5a9bb9cd1d82..d0b7037700e8 100644 --- a/tests/gateway/test_discord_reply_mode.py +++ b/tests/gateway/test_discord_reply_mode.py @@ -230,7 +230,7 @@ def test_from_dict_defaults_to_first(self): class TestEnvVarOverride: - """Tests for DISCORD_REPLY_TO_MODE environment variable override.""" + """Tests for Discord environment variable overrides.""" def _make_config(self): config = GatewayConfig() @@ -275,3 +275,27 @@ def test_env_var_creates_platform_config_if_missing(self): _apply_env_overrides(config) assert Platform.DISCORD in config.platforms assert config.platforms[Platform.DISCORD].reply_to_mode == "off" + + def test_allow_bots_env_override_saved_to_extra(self): + config = self._make_config() + with patch.dict(os.environ, {"DISCORD_ALLOW_BOTS": "mentions"}, clear=False): + _apply_env_overrides(config) + assert config.platforms[Platform.DISCORD].extra["allow_bots"] == "mentions" + + def test_target_bot_ids_env_override_saved_to_extra(self): + config = self._make_config() + with patch.dict(os.environ, {"DISCORD_TARGET_BOT_IDS": "111, 222"}, clear=False): + _apply_env_overrides(config) + assert config.platforms[Platform.DISCORD].extra["target_bot_ids"] == ["111", "222"] + + def test_auto_mention_target_bots_env_override_saved_to_extra(self): + config = self._make_config() + with patch.dict(os.environ, {"DISCORD_AUTO_MENTION_TARGET_BOTS": "true"}, clear=False): + _apply_env_overrides(config) + assert config.platforms[Platform.DISCORD].extra["auto_mention_target_bots"] is True + + def test_accept_bot_replies_env_override_saved_to_extra(self): + config = self._make_config() + with patch.dict(os.environ, {"DISCORD_ACCEPT_BOT_REPLIES": "false"}, clear=False): + _apply_env_overrides(config) + assert config.platforms[Platform.DISCORD].extra["accept_bot_replies"] is False diff --git a/tests/gateway/test_discord_send.py b/tests/gateway/test_discord_send.py index 8883d46efc2c..fecd06d6c662 100644 --- a/tests/gateway/test_discord_send.py +++ b/tests/gateway/test_discord_send.py @@ -78,3 +78,69 @@ async def fake_send(*, content, reference=None): assert channel.send.await_count == 2 assert send_calls[0]["reference"] is ref_msg assert send_calls[1]["reference"] is None + + +@pytest.mark.asyncio +async def test_send_prefixes_explicit_bot_mentions_from_metadata(): + adapter = DiscordAdapter(PlatformConfig(enabled=True, token="***")) + + sent_msg = SimpleNamespace(id=456) + send_calls = [] + + async def fake_send(**kwargs): + send_calls.append(kwargs) + return sent_msg + + channel = SimpleNamespace( + fetch_message=AsyncMock(), + send=AsyncMock(side_effect=fake_send), + ) + adapter._client = SimpleNamespace( + get_channel=lambda _chat_id: channel, + fetch_channel=AsyncMock(), + ) + + result = await adapter.send( + "555", + "hello there", + metadata={"mention_bot_ids": ["123456789"]}, + ) + + assert result.success is True + assert send_calls[0]["content"].startswith("<@123456789> hello there") + assert "allowed_mentions" in send_calls[0] + + +@pytest.mark.asyncio +async def test_send_can_auto_mention_configured_target_bots(): + adapter = DiscordAdapter( + PlatformConfig( + enabled=True, + token="***", + extra={ + "target_bot_ids": ["111", "222"], + "auto_mention_target_bots": True, + }, + ) + ) + + sent_msg = SimpleNamespace(id=789) + send_calls = [] + + async def fake_send(**kwargs): + send_calls.append(kwargs) + return sent_msg + + channel = SimpleNamespace( + fetch_message=AsyncMock(), + send=AsyncMock(side_effect=fake_send), + ) + adapter._client = SimpleNamespace( + get_channel=lambda _chat_id: channel, + fetch_channel=AsyncMock(), + ) + + result = await adapter.send("555", "ping") + + assert result.success is True + assert send_calls[0]["content"].startswith("<@111> <@222> ping")