fix(gateway): isolate /voice state and voice-channel input per multiplexed profile (#84872)

Voice state was keyed `<platform>:<chat_id>` with no profile namespace,
so two bots in one Discord channel shared one /voice mode; every
`_voice_input_callback` was the bare `_handle_voice_channel_input`, which
(like `_handle_voice_timeout_cleanup` and the /voice slash handler) always
picked `self.adapters[DISCORD]` — a secondary profile's voice transcripts
were dispatched through the default profile's bot.

- `_voice_key(platform, chat_id, profile=None)`: named profiles get a
  `<profile>:` prefix; default keeps the legacy shape (persisted state valid).
- `_voice_key_for_source` keys by the transport-OWNING profile
  (`_adapter_profile_for_source`), matching what `_sync_voice_mode_state_to_adapter`
  now restores per adapter via `_owner_profile`.
- `_bind_voice_input_callback` binds the capturing adapter into the
  transcript handler (functools.partial); used at primary connect,
  primary reconnect, /voice channel join, and `_configure_profile_adapter`.
- `_handle_voice_timeout_cleanup` takes the adapter it was bound to.
- /voice, join, leave and `_should_send_voice_reply` resolve the adapter via
  `_adapter_for_source` (fail-closed) instead of `self.adapters[platform]`.
- #84872: `_start_one_profile_adapters` now calls
  `_sync_voice_mode_state_to_adapter` on secondary INITIAL connect, as the
  primary path and both reconnect paths already did.

Co-authored-by: davidxyuan <124700534+davidxyuan@users.noreply.github.com>
This commit is contained in:
Teknium
2026-09-02 04:01:12 -07:00
parent e8231f01da
commit ca42d7a034
4 changed files with 184 additions and 26 deletions
+73 -21
View File
@@ -28,6 +28,7 @@ import asyncio
import concurrent.futures
import dataclasses
import faulthandler
import functools
import inspect
import json
import logging
@@ -8168,9 +8169,43 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
_VOICE_MODE_PATH = _hermes_home / "gateway_voice_mode.json"
def _voice_key(self, platform: Platform, chat_id: str) -> str:
"""Return a platform-namespaced key for voice mode state."""
return f"{platform.value}:{chat_id}"
def _voice_key(
self, platform: Platform, chat_id: str, profile: Optional[str] = None
) -> str:
"""Return a platform-namespaced key for voice mode state.
Under multiplexing the key is additionally namespaced by the profile
whose bot speaks in the chat (``<profile>:<platform>:<chat_id>``); the
default profile keeps the historical ``<platform>:<chat_id>`` shape so
persisted state stays valid. Two bots in one Discord channel otherwise
share a key and one profile's ``/voice`` flips the other's (#75198).
"""
base = f"{platform.value}:{chat_id}"
profile = profile.strip() if isinstance(profile, str) else ""
if not profile or profile == "default":
return base
return f"{profile}:{base}"
def _voice_key_for_source(self, source: SessionSource) -> str:
"""Voice-state key for an inbound source, namespaced by its transport owner.
Voice mode belongs to the (bot, chat) pair, so the namespace is the
profile that OWNS the receiving adapter (``_adapter_profile_for_source``)
— the same profile ``_sync_voice_mode_state_to_adapter`` uses on
reconnect — not the routed runtime profile.
"""
return self._voice_key(
source.platform,
source.chat_id,
profile=self._adapter_profile_for_source(source),
)
def _bind_voice_input_callback(self, adapter) -> None:
"""Route voice transcripts back through the adapter that captured them."""
if hasattr(adapter, "_voice_input_callback"):
adapter._voice_input_callback = functools.partial(
self._handle_voice_channel_input, adapter=adapter
)
def _load_voice_modes(self) -> Dict[str, str]:
try:
@@ -8269,7 +8304,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
if hasattr(adapter, "_auto_tts_default"):
adapter._auto_tts_default = _auto_tts_default
prefix = f"{platform.value}:"
prefix = self._voice_key(platform, "", profile=getattr(adapter, "_owner_profile", None))
if isinstance(disabled_chats, set):
disabled_chats.clear()
disabled_chats.update(
@@ -14276,8 +14311,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
self._sync_voice_mode_state_to_adapter(adapter)
# Wire voice input callback at connect time so voice
# transcription is forwarded without requiring /voice join.
if hasattr(adapter, "_voice_input_callback"):
adapter._voice_input_callback = self._handle_voice_channel_input
self._bind_voice_input_callback(adapter)
connected_count += 1
self._update_platform_runtime_status(
platform.value, platform_state="connected", error_code=None, error_message=None,
@@ -16040,8 +16074,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
self.adapters[platform] = adapter
self._sync_voice_mode_state_to_adapter(adapter)
# Wire voice input callback on reconnect as well (#60623).
if hasattr(adapter, "_voice_input_callback"):
adapter._voice_input_callback = self._handle_voice_channel_input
self._bind_voice_input_callback(adapter)
self.delivery_router.adapters = self.adapters
del self._failed_platforms[platform]
self._update_platform_runtime_status(
@@ -17156,6 +17189,9 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
)
if success:
profile_map[platform] = adapter
# Restore persisted /voice state for this bot (#84872) —
# primary startup and every reconnect path already do.
self._sync_voice_mode_state_to_adapter(adapter)
if credential_claim is not None:
claimed[credential_claim] = profile_name
if listener_claim is not None:
@@ -17214,6 +17250,9 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
adapter.set_platform_event_handler(
self._make_profile_platform_event_handler(profile_name)
)
# Voice transcripts from this bot's channels dispatch through THIS
# adapter (primary wiring lives at connect time; see #75198).
self._bind_voice_input_callback(adapter)
text_modes = getattr(self, "_busy_text_modes_by_profile", None)
adapter._busy_text_mode = (
text_modes.get(profile_name, self._busy_text_mode)
@@ -24549,15 +24588,18 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
# Wire callbacks BEFORE join so voice input arriving immediately
# after connection is not lost.
if hasattr(adapter, "_voice_input_callback"):
adapter._voice_input_callback = self._handle_voice_channel_input
self._bind_voice_input_callback(adapter)
voice_profile = self._adapter_profile_for_source(event.source)
if hasattr(adapter, "_on_voice_disconnect"):
adapter._on_voice_disconnect = self._handle_voice_timeout_cleanup
adapter._on_voice_disconnect = functools.partial(
self._handle_voice_timeout_cleanup, adapter=adapter
)
# Let the adapter's inactivity timer see the live voice-reply mode so it
# doesn't disconnect a deliberately text-only (/voice off) session.
if hasattr(adapter, "_voice_mode_getter"):
adapter._voice_mode_getter = lambda chat_id: self._voice_mode.get(
self._voice_key(Platform.DISCORD, str(chat_id)), "off"
self._voice_key(Platform.DISCORD, str(chat_id), profile=voice_profile),
"off",
)
try:
@@ -24577,7 +24619,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
adapter._voice_text_channels[guild_id] = int(event.source.chat_id)
if hasattr(adapter, "_voice_sources"):
adapter._voice_sources[guild_id] = event.source.to_dict()
self._voice_mode[self._voice_key(event.source.platform, event.source.chat_id)] = "all"
self._voice_mode[self._voice_key_for_source(event.source)] = "all"
self._save_voice_modes()
self._set_adapter_auto_tts_enabled(adapter, event.source.chat_id, enabled=True)
return (
@@ -24604,21 +24646,26 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
except Exception as e:
logger.warning("Error leaving voice channel: %s", e)
# Always clean up state even if leave raised an exception
self._voice_mode[self._voice_key(event.source.platform, event.source.chat_id)] = "off"
self._voice_mode[self._voice_key_for_source(event.source)] = "off"
self._save_voice_modes()
self._set_adapter_auto_tts_disabled(adapter, event.source.chat_id, disabled=True)
if hasattr(adapter, "_voice_input_callback"):
adapter._voice_input_callback = None
return "Left voice channel."
def _handle_voice_timeout_cleanup(self, chat_id: str) -> None:
def _handle_voice_timeout_cleanup(self, chat_id: str, *, adapter=None) -> None:
"""Called by the adapter when a voice channel times out.
Cleans up runner-side voice_mode state that the adapter cannot reach.
``adapter`` is the Discord adapter that timed out (bound at join time);
under multiplexing that is a specific profile's bot, not necessarily
``self.adapters[DISCORD]``.
"""
self._voice_mode[self._voice_key(Platform.DISCORD, chat_id)] = "off"
if adapter is None:
adapter = self.adapters.get(Platform.DISCORD)
profile = getattr(adapter, "_owner_profile", None)
self._voice_mode[self._voice_key(Platform.DISCORD, chat_id, profile=profile)] = "off"
self._save_voice_modes()
adapter = self.adapters.get(Platform.DISCORD)
self._set_adapter_auto_tts_disabled(adapter, chat_id, disabled=True)
def _is_duplicate_voice_transcript(self, guild_id: int, user_id: int, transcript: str) -> bool:
@@ -24663,14 +24710,18 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
return False
async def _handle_voice_channel_input(
self, guild_id: int, user_id: int, transcript: str
self, guild_id: int, user_id: int, transcript: str, *, adapter=None
):
"""Handle transcribed voice from a user in a voice channel.
Creates a synthetic MessageEvent and processes it through the
adapter's full message pipeline (session, typing, agent, TTS reply).
``adapter`` is the Discord adapter that captured the audio (bound via
``_bind_voice_input_callback``); under multiplexing each profile's bot
must dispatch through its own adapter, never the default profile's.
"""
adapter = self.adapters.get(Platform.DISCORD)
if adapter is None:
adapter = self.adapters.get(Platform.DISCORD)
if not adapter:
return
@@ -24692,6 +24743,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
user_id=str(user_id),
user_name=str(user_id),
chat_type="channel",
profile=getattr(adapter, "_owner_profile", None),
)
# Check authorization before processing voice input
@@ -24763,11 +24815,11 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
return False
chat_id = event.source.chat_id
voice_key = self._voice_key(event.source.platform, chat_id)
voice_key = self._voice_key_for_source(event.source)
voice_mode = self._voice_mode.get(voice_key)
is_voice_input = (event.message_type == MessageType.VOICE)
adapter = self.adapters.get(event.source.platform)
adapter = self._adapter_for_source(event.source)
adapter_auto_tts = False
if adapter and hasattr(adapter, "_should_auto_tts_for_chat"):
try:
+5 -4
View File
@@ -3347,10 +3347,12 @@ class GatewaySlashCommandsMixin:
"""Handle /voice [on|off|tts|channel|leave|status] command."""
args = event.get_command_args().strip().lower()
chat_id = event.source.chat_id
platform = event.source.platform
voice_key = self._voice_key(platform, chat_id)
# Voice state belongs to the (bot, chat) pair: resolve the adapter that
# received the command and key the mode by its owning profile so two
# multiplexed bots in one chat keep independent /voice state (#75198).
voice_key = self._voice_key_for_source(event.source)
adapter = self.adapters.get(platform)
adapter = self._adapter_for_source(event.source)
if args in {"on", "enable"}:
self._voice_mode[voice_key] = "voice_only"
@@ -3382,7 +3384,6 @@ class GatewaySlashCommandsMixin:
"all": t("gateway.voice.label_all"),
}
# Append voice channel info if connected
adapter = self.adapters.get(event.source.platform)
guild_id = self._get_guild_id(event)
if guild_id and hasattr(adapter, "get_voice_channel_info"):
info = adapter.get_voice_channel_info(guild_id)
@@ -325,6 +325,27 @@ class TestSecondaryProfileFatalRecovery:
assert hydration_flags and set(hydration_flags) == {False}
assert runner._profile_adapters["reviewer"][Platform.DISCORD] is replacement
@pytest.mark.asyncio
async def test_secondary_initial_connect_syncs_voice_mode_state(self, monkeypatch):
"""#84872: a secondary bot gets its persisted /voice state at INITIAL
connect, not only on reconnect."""
runner = _secondary_recovery_runner()
adapter = _SecondaryRecoveryAdapter()
_install_secondary_reconnect_context(monkeypatch, runner, adapter)
synced = []
runner._sync_voice_mode_state_to_adapter = synced.append
monkeypatch.setattr("hermes_cli.env_loader.hydrate_profile_secret_sources", lambda h: {})
monkeypatch.setattr(gateway_run, "_load_gateway_runtime_config", lambda: {})
monkeypatch.setattr(runner, "_snapshot_profile_busy_modes", lambda *a, **k: None)
monkeypatch.setattr("hermes_cli.plugins.discover_plugins", lambda: None)
async def connect(a, platform):
return True
monkeypatch.setattr(runner, "_connect_initial_adapter_with_timeout", connect)
assert await runner._start_one_profile_adapters("reviewer", Path("/profiles/reviewer"), {}) == 1
assert synced == [adapter]
@pytest.mark.asyncio
async def test_retryable_secondary_fatal_reconnects_with_its_profile_scope(
self, monkeypatch
@@ -9,7 +9,9 @@ same key. The fix prefixes keys with platform value: 'telegram:123' vs
import json
import tempfile
from pathlib import Path
from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from gateway.config import Platform
@@ -115,6 +117,88 @@ class TestSyncVoiceModeStateToAdapter:
assert mock_adapter._auto_tts_disabled_chats == {"123"}
class TestVoiceModeProfileIsolation:
"""Two multiplexed bots in one Discord channel keep independent /voice
state and voice transcripts dispatch through the bot that heard them
(#75198 voice half)."""
@staticmethod
def _discord_adapter(owner=None):
from unittest.mock import AsyncMock
a = MagicMock()
a.platform = Platform.DISCORD
a._owner_profile = owner
a._voice_text_channels = {111: 123}
a._voice_sources = {}
a._voice_input_callback = None
a._on_voice_disconnect = None
a._voice_mode_getter = None
a._auto_tts_enabled_chats = set()
a._auto_tts_disabled_chats = set()
a._client = MagicMock()
a._client.get_channel = MagicMock(return_value=None)
a.handle_message = AsyncMock()
return a
@pytest.mark.asyncio
async def test_voice_state_and_transcripts_stay_with_the_owning_bot(self, tmp_path):
from types import SimpleNamespace
from gateway.platforms.base import MessageEvent, MessageType, SessionSource
runner = _make_runner()
runner._VOICE_MODE_PATH = tmp_path / "voice.json"
runner._is_user_authorized = lambda source: True
default_ad = self._discord_adapter()
bot2_ad = self._discord_adapter(owner="bot2")
runner.adapters = {Platform.DISCORD: default_ad}
runner._profile_adapters = {"bot2": {Platform.DISCORD: bot2_ad}}
# Inbound event from bot2's transport in channel 123 (same id the
# default bot also sees).
src = SessionSource(platform=Platform.DISCORD, chat_id="123", user_id="u1",
chat_type="channel", profile="bot2")
src._transport_adapter_ref = lambda: bot2_ad
await runner._handle_voice_command(
MessageEvent(text="/voice tts", message_type=MessageType.TEXT, source=src)
)
assert runner._voice_mode == {"bot2:discord:123": "all"}
assert "123" in bot2_ad._auto_tts_enabled_chats
assert "123" not in default_ad._auto_tts_enabled_chats
# A transcript captured by bot2's adapter runs through bot2, not default.
runner._bind_voice_input_callback(bot2_ad)
await bot2_ad._voice_input_callback(guild_id=111, user_id=42, transcript="hi")
bot2_ad.handle_message.assert_awaited_once()
default_ad.handle_message.assert_not_awaited()
assert bot2_ad.handle_message.call_args[0][0].source.profile == "bot2"
# Timeout cleanup from bot2's channel disables bot2's auto-TTS only.
join = MessageEvent(text="/voice channel", message_type=MessageType.TEXT, source=src)
join.raw_message = SimpleNamespace(guild_id=111, guild=None)
bot2_ad.join_voice_channel = AsyncMock(return_value=True)
ch = MagicMock(); ch.name = "General"
bot2_ad.get_user_voice_channel = AsyncMock(return_value=ch)
await runner._handle_voice_channel_join(join)
bot2_ad._on_voice_disconnect("123")
assert runner._voice_mode["bot2:discord:123"] == "off"
assert "123" in bot2_ad._auto_tts_disabled_chats
assert "123" not in default_ad._auto_tts_disabled_chats
def test_sync_restores_only_the_owning_profiles_chats(self):
runner = _make_runner()
runner._voice_mode = {"discord:1": "all", "bot2:discord:2": "all"}
default_ad = MagicMock(); default_ad.platform = Platform.DISCORD
default_ad._owner_profile = None; default_ad._auto_tts_enabled_chats = set()
bot2_ad = MagicMock(); bot2_ad.platform = Platform.DISCORD
bot2_ad._owner_profile = "bot2"; bot2_ad._auto_tts_enabled_chats = set()
runner._sync_voice_mode_state_to_adapter(default_ad)
runner._sync_voice_mode_state_to_adapter(bot2_ad)
assert default_ad._auto_tts_enabled_chats == {"1"}
assert bot2_ad._auto_tts_enabled_chats == {"2"}
# ---------------------------------------------------------------------------
# Helper
# ---------------------------------------------------------------------------