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:
+73
-21
@@ -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:
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user