From 0c58b7a9e4a16852edad74e26c3971abd0fb0c46 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 20:26:41 -0700 Subject: [PATCH] =?UTF-8?q?refactor(gateway):=20voice/watchers/tts/lease?= =?UTF-8?q?=20=E2=80=94=20phase=20helpers,=20folded=20probes,=20compact=20?= =?UTF-8?q?docstrings?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gateway/run_voice.py | 241 +++++++++++++++--------------- gateway/run_watchers.py | 156 ++++++++++--------- gateway/streaming_tts_consumer.py | 114 +++++++------- gateway/turn_lease.py | 45 +++--- 4 files changed, 265 insertions(+), 291 deletions(-) diff --git a/gateway/run_voice.py b/gateway/run_voice.py index 535441ed31..e814784244 100644 --- a/gateway/run_voice.py +++ b/gateway/run_voice.py @@ -16,6 +16,8 @@ import re import sys import time from contextlib import suppress +from difflib import SequenceMatcher +from types import SimpleNamespace from typing import TYPE_CHECKING, Dict, List, Optional from gateway.config import Platform @@ -30,49 +32,47 @@ logger = logging.getLogger("gateway.run") # Adapter-side per-chat auto-TTS override sets (``/voice off`` vs explicit ``/voice on``/``tts``). _OFF_SET, _ON_SET = "_auto_tts_disabled_chats", "_auto_tts_enabled_chats" +_VOICE_MODES = {"off", "voice_only", "all"} class GatewayVoiceMixin: """Voice-channel / auto-TTS methods for GatewayRunner.""" def _voice_key(self, platform: Platform, chat_id: str, profile: Optional[str] = None) -> str: - """Voice-mode state key: ``::`` under multiplexing (profile - whose bot speaks — else two bots in one channel share a key and one ``/voice`` flips the - other's); the default profile keeps ``:`` so persisted state stays valid. - """ + """``::`` under multiplexing (profile whose bot speaks — else + two bots in one channel share a key and one ``/voice`` flips the other's); the default + profile keeps ``:`` so persisted state stays valid.""" base = f"{platform.value}:{chat_id}" profile = profile.strip() if isinstance(profile, str) else "" return base if not profile or profile == "default" else f"{profile}:{base}" def _voice_key_for_source(self, source: SessionSource) -> str: - """Voice-state key for an inbound source. Voice mode belongs to the (bot, chat) pair, so the - namespace is the profile that OWNS the receiving adapter, not the routed profile.""" + """Voice mode belongs to the (bot, chat) pair: namespace is the profile that OWNS the + receiving adapter, not the routed profile.""" profile = self._adapter_profile_for_source(source) return self._voice_key(source.platform, source.chat_id, profile=profile) def _bind_voice_input_callback(self, adapter) -> None: """Route voice transcripts back through the adapter that captured them.""" if hasattr(adapter, "_voice_input_callback"): - cb = functools.partial(self._handle_voice_channel_input, adapter=adapter) - adapter._voice_input_callback = cb + adapter._voice_input_callback = functools.partial( + self._handle_voice_channel_input, adapter=adapter + ) def _load_voice_modes(self) -> Dict[str, str]: try: data = json.loads(self._VOICE_MODE_PATH.read_text(encoding="utf-8")) except (FileNotFoundError, json.JSONDecodeError, OSError): return {} - result = {} - for chat_id, mode in (data.items() if isinstance(data, dict) else ()): - if mode not in {"off", "voice_only", "all"}: - continue - if ":" in str(chat_id): - result[str(chat_id)] = mode - else: # legacy unprefixed key: warn and skip - logger.warning( - "Skipping legacy unprefixed voice mode key %r during migration. " - "Re-enable voice mode on that chat to rebuild the prefixed key.", str(chat_id), - ) - return result + if not isinstance(data, dict): + return {} + items = {str(k): m for k, m in data.items() if m in _VOICE_MODES} + for chat_id in (k for k in items if ":" not in k): # legacy unprefixed key: warn and skip + logger.warning( + "Skipping legacy unprefixed voice mode key %r during migration. " + "Re-enable voice mode on that chat to rebuild the prefixed key.", chat_id, + ) + return {k: m for k, m in items.items() if ":" in k} def _save_voice_modes(self) -> None: try: @@ -84,20 +84,19 @@ class GatewayVoiceMixin: @staticmethod def _toggle_adapter_auto_tts_set(adapter, chat_id: str, on: bool, *, enable: bool) -> None: - """Add/discard ``chat_id`` in the adapter's enabled (``enable=True``) or disabled set; adding - also clears the other set (``/voice off`` and ``/voice on``/``tts`` override each other). - """ + """Add/discard ``chat_id`` in the adapter's enabled (``enable=True``) or disabled set; + adding also clears the other set (``/voice off`` and ``/voice on``/``tts`` override each + other).""" add_to, clear_from = (_ON_SET, _OFF_SET) if enable else (_OFF_SET, _ON_SET) target = getattr(adapter, add_to, None) - other = getattr(adapter, clear_from, None) if not isinstance(target, set): return if not on: target.discard(chat_id) - else: - target.add(chat_id) - if isinstance(other, set): - other.discard(chat_id) + return + target.add(chat_id) + if isinstance(other := getattr(adapter, clear_from, None), set): + other.discard(chat_id) def _set_adapter_auto_tts_disabled(self, adapter, chat_id: str, disabled: bool) -> None: """Update an adapter's in-memory auto-TTS suppression set if present.""" @@ -108,18 +107,17 @@ class GatewayVoiceMixin: self._toggle_adapter_auto_tts_set(adapter, chat_id, enabled, enable=True) def _sync_voice_mode_state_to_adapter(self, adapter) -> None: - """Restore persisted /voice state into a live platform adapter: ``_auto_tts_default`` from - ``voice.auto_tts``; ``_auto_tts_enabled_chats`` (modes ``voice_only``/``all``) and - ``_auto_tts_disabled_chats`` (mode ``off``) from ``self._voice_mode``. - """ + """Restore persisted /voice state into a live adapter: ``_auto_tts_default`` from + ``voice.auto_tts``; enabled (``voice_only``/``all``) and disabled (``off``) chat sets from + ``self._voice_mode``.""" platform = getattr(adapter, "platform", None) if not isinstance(platform, Platform): return chat_sets = [ - (getattr(adapter, name, None), modes) + (chats, modes) for name, modes in ((_OFF_SET, {"off"}), (_ON_SET, {"voice_only", "all"})) + if isinstance(chats := getattr(adapter, name, None), set) ] - chat_sets = [(chats, modes) for chats, modes in chat_sets if isinstance(chats, set)] if not chat_sets: return # Lazy import: no module-level dep from gateway -> hermes_cli. @@ -144,9 +142,7 @@ class GatewayVoiceMixin: raw = getattr(event, "raw_message", None) if getattr(raw, "guild_id", None): # slash command interaction return int(raw.guild_id) - if getattr(raw, "guild", None): # regular message - return raw.guild.id - return None + return raw.guild.id if getattr(raw, "guild", None) else None # regular message async def _handle_voice_channel_join(self, event: MessageEvent) -> str: """Join the user's current Discord voice channel.""" @@ -177,12 +173,10 @@ class GatewayVoiceMixin: except Exception as e: logger.warning("Failed to join voice channel: %s", e) adapter._voice_input_callback = None - if any(tok in str(e).lower() for tok in ("pynacl", "nacl", "davey")): - return ( - "Voice dependencies are missing (PyNaCl / davey). " - f"Install with: `{sys.executable} -m pip install PyNaCl`" - ) - return f"Failed to join voice channel: {e}" + if not any(tok in str(e).lower() for tok in ("pynacl", "nacl", "davey")): + return f"Failed to join voice channel: {e}" + return ("Voice dependencies are missing (PyNaCl / davey). " + f"Install with: `{sys.executable} -m pip install PyNaCl`") if not success: adapter._voice_input_callback = None return "Failed to join voice channel. Check bot permissions (Connect + Speak)." @@ -200,11 +194,9 @@ class GatewayVoiceMixin: """Leave the Discord voice channel.""" adapter = self._adapter_for_source(event.source) guild_id = self._get_guild_id(event) - if ( - not guild_id - or not hasattr(adapter, "leave_voice_channel") - or not hasattr(adapter, "is_in_voice_channel") - or not adapter.is_in_voice_channel(guild_id) + if not ( + guild_id and hasattr(adapter, "leave_voice_channel") + and hasattr(adapter, "is_in_voice_channel") and adapter.is_in_voice_channel(guild_id) ): return "Not in a voice channel." try: @@ -220,10 +212,7 @@ class GatewayVoiceMixin: def _handle_voice_timeout_cleanup(self, chat_id: str, *, adapter=None) -> None: """Adapter callback on voice-channel timeout: clear runner-side voice_mode state. - - ``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]``. - """ + ``adapter`` (bound at join) is that profile's bot, not always ``self.adapters[DISCORD]``.""" if adapter is None: adapter = self.adapters.get(Platform.DISCORD) profile = getattr(adapter, "_owner_profile", None) @@ -239,8 +228,6 @@ class GatewayVoiceMixin: """Suppress repeated STT outputs for the same recent utterance (voice capture can emit an utterance twice a few seconds apart -> a second queued run and overlapping spoken replies). """ - from difflib import SequenceMatcher - normalized = re.sub(r"[^\w\s]", "", re.sub(r"\s+", " ", transcript).strip().lower()) if not normalized: return False @@ -250,43 +237,53 @@ class GatewayVoiceMixin: if not isinstance(recent_store, dict): recent_store = self._recent_voice_transcripts = {} recent = [(ts, txt) for ts, txt in recent_store.get(key, []) if now - ts <= 12.0] - for _, prior in recent: - if prior == normalized or ( + if any( + prior == normalized or ( len(prior) >= 16 and len(normalized) >= 16 and SequenceMatcher(None, prior, normalized).ratio() >= 0.95 - ): - recent_store[key] = recent - return True - recent.append((now, normalized)) - recent_store[key] = recent[-5:] + ) + for _, prior in recent + ): + recent_store[key] = recent + return True + recent_store[key] = (recent + [(now, normalized)])[-5:] return False + @staticmethod + def _voice_input_source(adapter, guild_id: int, user_id: int, text_ch_id) -> SessionSource: + """Bound text channel's own source when available (voice shares the text conversation's + session), else a synthetic one.""" + if source_data := getattr(adapter, "_voice_sources", {}).get(guild_id): + source = SessionSource.from_dict(source_data) + source.user_id = source.user_name = str(user_id) + return source + return SessionSource( + platform=Platform.DISCORD, chat_id=str(text_ch_id), user_id=str(user_id), + user_name=str(user_id), chat_type="channel", + profile=getattr(adapter, "_owner_profile", None), + ) + + @staticmethod + def _voice_channel_prompt(adapter, text_ch_id) -> Optional[str]: + """Bound text channel's channel_prompt: voice input gets the same per-channel context.""" + if callable(resolver := getattr(adapter, "_resolve_channel_prompt", None)): + with suppress(Exception): + resolved = resolver(str(text_ch_id)) + return resolved if isinstance(resolved, str) else None + return None + async def _handle_voice_channel_input( self, guild_id: int, user_id: int, transcript: str, *, adapter=None ): - """Handle transcribed voice from a user in a voice channel. - - ``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. - """ + """Handle transcribed voice from a voice channel. ``adapter`` captured the audio (bound + via ``_bind_voice_input_callback``); under multiplexing each profile's bot dispatches + through its own adapter, never the default profile's.""" if adapter is None: adapter = self.adapters.get(Platform.DISCORD) text_ch_id = adapter._voice_text_channels.get(guild_id) if adapter else None if not text_ch_id: return - # Reuse the linked text channel's source metadata when available so voice input shares - # the same session as the bound text conversation. - source_data = getattr(adapter, "_voice_sources", {}).get(guild_id) - if source_data: - source = SessionSource.from_dict(source_data) - source.user_id = source.user_name = str(user_id) - else: - source = SessionSource( - platform=Platform.DISCORD, chat_id=str(text_ch_id), user_id=str(user_id), - user_name=str(user_id), chat_type="channel", - profile=getattr(adapter, "_owner_profile", None), - ) + source = self._voice_input_source(adapter, guild_id, user_id, text_ch_id) if not self._is_user_authorized(source): logger.debug("Unauthorized voice input from user %d, ignoring", user_id) return @@ -302,34 +299,22 @@ class GatewayVoiceMixin: if channel: safe_text = transcript[:2000].replace("@everyone", "@\u200beveryone").replace("@here", "@\u200bhere") await channel.send(f"**[Voice]** <@{user_id}>: {safe_text}") - # Bound text channel's channel_prompt, so voice input gets the same per-channel context - # as typed messages. - channel_prompt = None - resolver = getattr(adapter, "_resolve_channel_prompt", None) - if callable(resolver): - with suppress(Exception): - resolved = resolver(str(text_ch_id)) - channel_prompt = resolved if isinstance(resolved, str) else None # Synthetic MessageEvent for the normal pipeline; the SimpleNamespace raw_message lets # _get_guild_id() extract guild_id so _send_voice_reply() plays audio in the voice channel. - from types import SimpleNamespace event = MessageEvent( source=source, text=transcript, message_type=MessageType.VOICE, raw_message=SimpleNamespace(guild_id=guild_id, guild=None), - channel_prompt=channel_prompt, + channel_prompt=self._voice_channel_prompt(adapter, text_ch_id), ) await adapter.handle_message(event) def _should_send_voice_reply( self, event: MessageEvent, response: str, agent_messages: list, already_sent: bool = False ) -> bool: - """Decide whether the runner should send a TTS voice reply. - - False when voice_mode is off for this chat, the response is empty/an error, the agent - already called text_to_speech this turn (dedup), or voice input + base adapter auto-TTS - already handled it — UNLESS streaming consumed the response (already_sent=True), since - then the base adapter has no text for auto-TTS and the runner must handle it. - """ + """False when voice_mode is off for this chat, the response is empty/an error, the agent + already called text_to_speech this turn, or voice input + base adapter auto-TTS handled + it — UNLESS streaming consumed the response (already_sent): then the base adapter has no + text for auto-TTS and the runner must handle it.""" if not response or response.startswith("Error:"): return False chat_id = event.source.chat_id @@ -353,14 +338,13 @@ class GatewayVoiceMixin: ) return False # Dedup: agent already called the TTS tool in THIS turn (from the last user message on). - turn_messages = agent_messages - for i in range(len(agent_messages) - 1, -1, -1): - if agent_messages[i].get("role") == "user": - turn_messages = agent_messages[i:] - break + start = next( + (i for i, m in reversed(list(enumerate(agent_messages))) if m.get("role") == "user"), + 0, + ) if any( (tc.get("function") or {}).get("name") == "text_to_speech" - for msg in turn_messages if msg.get("role") == "assistant" + for msg in agent_messages[start:] if msg.get("role") == "assistant" for tc in (msg.get("tool_calls") or []) ): return False @@ -373,12 +357,39 @@ class GatewayVoiceMixin: """Return whether inbound voice/STT transcripts should be echoed to chat.""" return bool(getattr(self.config, "stt_echo_transcripts", True)) + @staticmethod + async def _synthesize_voice_reply(text: str, audio_path: str) -> List[str]: + """Run the TTS tool for ``text`` into ``audio_path``; return the produced file paths (one + combined file, or several separately valid ones when combination is unavailable / over a + platform limit; legacy single-file results keep working) — ``[]`` on failure.""" + from tools.tts_tool import text_to_speech_tool + + result_json = await asyncio.to_thread( + text_to_speech_tool, text=text, output_path=audio_path + ) + try: + result = json.loads(result_json) + except (json.JSONDecodeError, TypeError): + logger.warning( + "Auto voice reply TTS returned invalid JSON: %s", + result_json[:200] if result_json else result_json, + ) + return [] + actual_paths = [ + str(p) for p in (result.get("file_paths") or [result.get("file_path", audio_path)]) + if p and os.path.isfile(p) + ] + if not result.get("success") or not actual_paths: + logger.warning("Auto voice reply TTS failed: %s", result.get("error")) + return [] + return actual_paths + async def _send_voice_reply(self, event: MessageEvent, text: str) -> None: """Generate TTS audio and send as a voice message before the text reply.""" audio_path = None actual_paths: List[str] = [] try: - from tools.tts_tool import text_to_speech_tool, _strip_markdown_for_tts + from tools.tts_tool import _strip_markdown_for_tts tts_text = _strip_markdown_for_tts(text) if not tts_text: @@ -386,24 +397,9 @@ class GatewayVoiceMixin: # Platforms whose native voice bubbles require Ogg/Opus (OPUS_VOICE_PLATFORMS) get an # explicit .ogg path; the TTS tool's container repair guarantees real Ogg/Opus bytes. audio_path = build_auto_tts_output_path(event.source.platform) - result_json = await asyncio.to_thread( - text_to_speech_tool, text=tts_text, output_path=audio_path - ) - try: - result = json.loads(result_json) - except (json.JSONDecodeError, TypeError): - logger.warning("Auto voice reply TTS returned invalid JSON: %s", result_json[:200] if result_json else result_json) - return - # One combined file or several separately valid files (combination unavailable or - # over a platform limit); legacy single-file results keep working. - actual_paths = [ - str(p) for p in (result.get("file_paths") or [result.get("file_path", audio_path)]) - if p and os.path.isfile(p) - ] - if not result.get("success") or not actual_paths: - logger.warning("Auto voice reply TTS failed: %s", result.get("error")) - return - await self._deliver_voice_reply(event, actual_paths) + actual_paths = await self._synthesize_voice_reply(tts_text, audio_path) + if actual_paths: + await self._deliver_voice_reply(event, actual_paths) except Exception as e: logger.warning("Auto voice reply failed: %s", e, exc_info=True) finally: @@ -417,12 +413,11 @@ class GatewayVoiceMixin: guild_id = self._get_guild_id(event) play = getattr(adapter, "play_in_voice_channel", None) is_in_vc = getattr(adapter, "is_in_voice_channel", None) - send_voice = getattr(adapter, "send_voice", None) if guild_id and callable(play) and callable(is_in_vc) and is_in_vc(guild_id): for path in audio_paths: await play(guild_id, path) return - if not callable(send_voice): + if not callable(send_voice := getattr(adapter, "send_voice", None)): return reply_anchor = self._reply_anchor_for_event(event) # Mark the auto voice reply notify-worthy (mirrors the final-text path in platforms/base.py) diff --git a/gateway/run_watchers.py b/gateway/run_watchers.py index 678813d9a1..79ba0d7533 100644 --- a/gateway/run_watchers.py +++ b/gateway/run_watchers.py @@ -41,15 +41,13 @@ class GatewaySessionWatchersMixin: """Session expiry / stall / catalog-refresh watcher loops for GatewayRunner.""" async def _session_expiry_watcher(self, interval: int = 300): - """Background task that finalizes expired sessions (``on_session_finalize`` hooks, cached - agent teardown, cache eviction, ``expiry_finalized`` flag) and runs the cache/store sweeps. - """ + """Finalize expired sessions (``on_session_finalize`` hooks, cached agent teardown, cache + eviction, ``expiry_finalized`` flag) and run the cache/store sweeps.""" await asyncio.sleep(60) # initial delay — let the gateway fully start finalize_failures: dict[str, int] = {} # session_id -> consecutive failure count while self._running: try: - expired = await self._collect_expired_sessions() - if expired: + if expired := await self._collect_expired_sessions(): platforms = Counter(_platform_of_key(k, "unknown") for k, _ in expired) logger.info( "Session expiry: %d sessions to finalize (%s)", @@ -70,24 +68,20 @@ class GatewaySessionWatchersMixin: async def _collect_expired_sessions(self) -> list: """Return ``[(session_key, entry)]`` for expired, not-yet-finalized sessions.""" - await self.async_session_store._ensure_loaded() - expired = [] - for key, entry in list(self.session_store._entries.items()): - if entry.expiry_finalized: - continue - if await self.async_session_store._is_session_expired(entry): - expired.append((key, entry)) - return expired + store = self.async_session_store + await store._ensure_loaded() + return [ + (key, entry) for key, entry in list(self.session_store._entries.items()) + if not entry.expiry_finalized and await store._is_session_expired(entry) + ] async def _finalize_expired_sessions(self, expired: list, failures: dict[str, int]) -> None: """Finalize each entry; after ``_MAX_FINALIZE_RETRIES`` consecutive failures mark it - finalized anyway (without clearing the model override) to stop an infinite retry loop. - """ + finalized anyway (without clearing the model override) to stop an infinite retry loop.""" for key, entry in expired: sid = entry.session_id try: await self._finalize_expired_session(key, entry) - failures.pop(sid, None) except Exception as e: count = failures[sid] = failures.get(sid, 0) + 1 if count < _MAX_FINALIZE_RETRIES: @@ -103,7 +97,20 @@ class GatewaySessionWatchersMixin: ) store = self.async_session_store await store.set_expiry_finalized(entry, clear_model_override=False) - failures.pop(sid, None) + failures.pop(sid, None) + + def _agent_for_expired_session(self, key: str): + """Idle agents live in _agent_cache (not _running_agents); fall back to the running turn's + agent in case the session is still mid-turn when the expiry fires.""" + cache_lock = getattr(self, "_agent_cache_lock", None) # tests build runners without it + if cache_lock is not None: + with cache_lock: + cached = self._agent_cache.get(key) + agent = cached[0] if isinstance(cached, tuple) else cached if cached else None + if agent is not None: + return agent + state = self._peek_session_state(key) + return state.turn.agent if state else None async def _finalize_expired_session(self, key: str, entry) -> None: """Run finalize hooks, tear down the cached agent, clear conversation scope, persist.""" @@ -119,17 +126,7 @@ class GatewaySessionWatchersMixin: ) except Exception: pass - # Idle agents live in _agent_cache (not _running_agents); fall back to the running - # turn's agent in case the session is still mid-turn when the expiry fires. - agent = None - cache_lock = getattr(self, "_agent_cache_lock", None) - if cache_lock is not None: - with cache_lock: - cached = self._agent_cache.get(key) - agent = cached[0] if isinstance(cached, tuple) else cached if cached else None - if agent is None: - state = self._peek_session_state(key) - agent = state.turn.agent if state else None + agent = self._agent_for_expired_session(key) if agent and agent is not _AGENT_PENDING_SENTINEL: await self._cleanup_agent_resources_off_loop(agent, context="session expiry") # Evict so the AIAgent (LLM clients, tool schemas, memory refs) can be GC'd, then drop every @@ -177,7 +174,7 @@ class GatewaySessionWatchersMixin: def _iter_gateway_adapters(self): """Yield every live platform adapter (default + multiplex profiles), deduped by identity.""" seen: set[int] = set() - maps = [getattr(self, "adapters", {}), *getattr(self, "_profile_adapters", {}).values()] + maps = (getattr(self, "adapters", {}), *getattr(self, "_profile_adapters", {}).values()) for amap in maps: for adapter in list(amap.values()): if adapter is not None and id(adapter) not in seen: @@ -185,9 +182,8 @@ class GatewaySessionWatchersMixin: yield adapter def _session_activity_for_stall(self, session_key: str) -> Optional[dict]: - """Return the shared activity snapshot for stall progress: the single source is - ``AIAgent.get_activity_summary()``; no turn-start or pending-inbound clocks. - """ + """Activity snapshot for stall progress: the single source is + ``AIAgent.get_activity_summary()``; no turn-start or pending-inbound clocks.""" from gateway.run import _AGENT_PENDING_SENTINEL agent = (getattr(self, "_running_agents", None) or {}).get(session_key) if agent is None or agent is _AGENT_PENDING_SENTINEL: @@ -226,8 +222,7 @@ class GatewaySessionWatchersMixin: async def _check_session_stalls(self, timeout_seconds: float) -> int: """Scan pending inbound sessions and notify once per stall episode; returns the number of - notifications sent this pass (for tests). - """ + notifications sent this pass (for tests).""" from gateway.session_stall import ( resolve_session_idle_seconds_from_activity, should_clear_session_stall_notification, @@ -237,9 +232,7 @@ class GatewaySessionWatchersMixin: notified_map = getattr(self, "_session_stall_notified", None) if notified_map is None: notified_map = self._session_stall_notified = {} - sent = 0 - now = time.time() - candidates = self._stall_candidates() + sent, now, candidates = 0, time.time(), self._stall_candidates() # Every candidate carries a non-None pending event, so has_pending_inbound is always True. for session_key, (adapter, pending_event) in list(candidates.items()): activity = self._session_activity_for_stall(session_key) @@ -266,46 +259,11 @@ class GatewaySessionWatchersMixin: notified_map.pop(key, None) return sent - async def _notify_session_stall( - self, session_key: str, adapter, pending_event, idle_seconds: float, activity: dict, - timeout_seconds: float, notified_map: dict, - ) -> bool: - """Log one stall episode and deliver the notice. True only when sent (latched); - undeliverable (no chat_id) latches without sending; send failures never latch.""" + async def _send_stall_notice(self, session_key: str, adapter, source, idle_seconds) -> bool: + """Deliver one stall notice, bounded and failure-tolerant; True only when delivered.""" from gateway.run import _STALL_NOTIFY_SEND_TIMEOUT_SECONDS - from gateway.session_stall import ( - format_session_stall_notification, - resolve_session_idle_seconds_from_activity, - ) + from gateway.session_stall import format_session_stall_notification - logger.warning( - "Session stall detected: session=%s idle=%.0fs (timeout=%.0fs, ~%d min); pending " - "inbound present | last_activity=%s | provenance=%s (agent.session_stall_timeout)", - session_key, idle_seconds, timeout_seconds, max(1, int(idle_seconds // 60)), - activity.get("last_activity_desc") or activity.get("last_activity_description") - or "unknown", - activity.get("provenance") or activity.get("last_activity_provenance") or "unknown", - ) - source = getattr(pending_event, "source", None) - if not (chat_id := getattr(source, "chat_id", None)): - logger.warning("Session stall notify skipped (no chat_id): session=%s", session_key) - notified_map[session_key] = True # cannot deliver; latch to avoid log spam every tick - return False - # Re-read pending state + activity IMMEDIATELY before delivery: the snapshot ages while - # earlier candidates await sends; an agent that progressed (or drained its queue) must not - # get a false stall notice. Abort with the latch un-set so the next tick re-evaluates. - still_pending = self._session_still_pending(adapter, session_key) - fresh_idle = resolve_session_idle_seconds_from_activity( - self._session_activity_for_stall(session_key), now=time.time() - ) - if not still_pending or (fresh_idle is not None and fresh_idle < timeout_seconds): - logger.info( - "Session stall notify aborted (no longer stale): " - "session=%s pending=%s fresh_idle=%s", - session_key, still_pending, fresh_idle, - ) - notified_map.pop(session_key, None) # re-arm so a FUTURE genuine stall notifies again - return False try: metadata = self._thread_metadata_for_source(source) notice = format_session_stall_notification(idle_seconds) @@ -313,7 +271,7 @@ class GatewaySessionWatchersMixin: # block the watcher pass — siblings would go unevaluated and the watcher stop. try: result = await asyncio.wait_for( - adapter.send(str(chat_id), notice, metadata=metadata), + adapter.send(str(source.chat_id), notice, metadata=metadata), timeout=_STALL_NOTIFY_SEND_TIMEOUT_SECONDS, ) except asyncio.TimeoutError: @@ -332,13 +290,52 @@ class GatewaySessionWatchersMixin: except Exception as exc: logger.warning("Session stall notify failed for %s: %s", session_key, exc) return False + return True + + async def _notify_session_stall( + self, session_key: str, adapter, pending_event, idle_seconds: float, activity: dict, + timeout_seconds: float, notified_map: dict, + ) -> bool: + """Log one stall episode and deliver the notice. True only when sent (latched); + undeliverable (no chat_id) latches without sending; send failures never latch.""" + from gateway.session_stall import resolve_session_idle_seconds_from_activity + + logger.warning( + "Session stall detected: session=%s idle=%.0fs (timeout=%.0fs, ~%d min); pending " + "inbound present | last_activity=%s | provenance=%s (agent.session_stall_timeout)", + session_key, idle_seconds, timeout_seconds, max(1, int(idle_seconds // 60)), + activity.get("last_activity_desc") or activity.get("last_activity_description") + or "unknown", + activity.get("provenance") or activity.get("last_activity_provenance") or "unknown", + ) + source = getattr(pending_event, "source", None) + if not getattr(source, "chat_id", None): + logger.warning("Session stall notify skipped (no chat_id): session=%s", session_key) + notified_map[session_key] = True # cannot deliver; latch to avoid log spam every tick + return False + # Re-read pending state + activity IMMEDIATELY before delivery: the snapshot ages while + # earlier candidates await sends; an agent that progressed (or drained its queue) must not + # get a false stall notice. Abort with the latch un-set so the next tick re-evaluates. + still_pending = self._session_still_pending(adapter, session_key) + fresh_idle = resolve_session_idle_seconds_from_activity( + self._session_activity_for_stall(session_key), now=time.time() + ) + if not still_pending or (fresh_idle is not None and fresh_idle < timeout_seconds): + logger.info( + "Session stall notify aborted (no longer stale): " + "session=%s pending=%s fresh_idle=%s", + session_key, still_pending, fresh_idle, + ) + notified_map.pop(session_key, None) # re-arm so a FUTURE genuine stall notifies again + return False + if not await self._send_stall_notice(session_key, adapter, source, idle_seconds): + return False notified_map[session_key] = True return True async def _model_catalog_refresh_watcher(self) -> None: """Refresh the /model picker's remote catalogs every TTL window. The picker itself only - refreshes on a cold/stale open, so if nobody opens ``/model`` the cache never updates. - """ + refreshes on a cold/stale open, so if nobody opens ``/model`` the cache never updates.""" from hermes_cli.model_catalog import refresh_catalogs, refresh_interval_seconds await asyncio.sleep(30) # let startup settle @@ -364,8 +361,7 @@ class GatewaySessionWatchersMixin: await asyncio.sleep(min(30.0, max(1.0, float(interval)))) while self._running: try: - timeout = self._session_stall_timeout_seconds() - if timeout > 0: + if (timeout := self._session_stall_timeout_seconds()) > 0: await self._check_session_stalls(timeout) except Exception as exc: logger.debug("Session stall watcher error: %s", exc) diff --git a/gateway/streaming_tts_consumer.py b/gateway/streaming_tts_consumer.py index 2895caea68..361be3535e 100644 --- a/gateway/streaming_tts_consumer.py +++ b/gateway/streaming_tts_consumer.py @@ -1,12 +1,10 @@ """Gateway streaming-TTS consumer — LLM deltas to adapter PCM audio sink. -Bridges the sync agent ``stream_delta_callback`` (worker thread) to a voice-capable adapter's -streaming-audio contract so playback begins while the LLM is still generating. ``on_delta`` -never blocks (SentenceChunker -> thread-safe queue); the ``_run`` task on the gateway loop -drains, synthesises via a ``StreamingTTSProvider`` and writes PCM. Outcome: full success -> -``completed``; failure before any audible output -> ``suppress_whole_file=False`` (gateway -falls back to whole-file TTS); failure after partial audio -> ``partial`` + suppress (never -replay the response from the beginning). +``on_delta`` (agent worker thread) never blocks: SentenceChunker -> thread-safe queue. The +``_run`` task on the gateway loop drains, synthesises via a ``StreamingTTSProvider`` and writes +PCM so playback starts mid-generation. Outcome: success -> ``completed``; failure before audible +output -> ``suppress_whole_file=False`` (gateway falls back to whole-file TTS); failure after +partial audio -> ``partial`` + suppress (never replay the response from the beginning). """ from __future__ import annotations @@ -41,10 +39,7 @@ class StreamingTTSConsumer: ) -> None: from tools.tts_streaming import SentenceChunker, resolve_streaming_provider - self._adapter = adapter - self._chat_id = chat_id - self._loop = loop - self._metadata = metadata + self._adapter, self._chat_id, self._loop, self._metadata = adapter, chat_id, loop, metadata # Resolved once; None => inactive, gateway falls back to whole-file TTS. self._streamer = resolve_streaming_provider(tts_config) self._chunker = SentenceChunker() @@ -64,20 +59,13 @@ class StreamingTTSConsumer: self._lock = threading.Lock() self._strip_markdown = None # lazily imported to avoid import cycles - # usable streaming provider resolved - active = property(lambda self: self._streamer is not None) - # streaming audio fully delivered - completed = property(lambda self: self._completed) - # some audio was audible before a failure/drop - partial = property(lambda self: self._partial) - # first PCM chunk has been written - audible = property(lambda self: bool(self._handle and self._handle.audible)) - # queue saturation dropped at least one clause - dropped = property(lambda self: self._dropped) - # gateway should skip whole-file TTS fallback - suppress_whole_file = property(lambda self: self._suppress_whole_file) - # async drain task has terminated - done = property(lambda self: self._task is not None and self._task.done()) + active = property(lambda self: self._streamer is not None) # usable streaming provider + completed = property(lambda self: self._completed) # streaming audio fully delivered + partial = property(lambda self: self._partial) # some audio audible before a failure/drop + audible = property(lambda self: bool(self._handle and self._handle.audible)) # PCM written + dropped = property(lambda self: self._dropped) # queue saturation dropped >= 1 clause + suppress_whole_file = property(lambda self: self._suppress_whole_file) # skip whole-file TTS + done = property(lambda self: self._task is not None and self._task.done()) # drain task ended def _enqueue_clauses(self, clauses, full_msg: str, *, log_errors: bool) -> None: try: @@ -100,8 +88,7 @@ class StreamingTTSConsumer: def finish(self) -> None: """Signal end-of-text, flush the chunker tail, then enqueue ``_DONE`` after all flushed - clauses so the drain loop terminates deterministically without racing a late ``on_delta``. - """ + clauses so the drain loop ends deterministically without racing a late ``on_delta``.""" if self._finished: return self._finished = True @@ -121,13 +108,13 @@ class StreamingTTSConsumer: self._queue.put_nowait(sentinel) return True except queue.Full: - try: - self._queue.get_nowait() - if mark_dropped: - self._dropped = True - except queue.Empty: - return not mark_dropped - return False + pass + try: + self._queue.get_nowait() + except queue.Empty: + return not mark_dropped + self._dropped = self._dropped or mark_dropped + return False def start(self) -> asyncio.Task: """Create (once) and return the async drain task on the gateway loop.""" @@ -138,20 +125,18 @@ class StreamingTTSConsumer: def _settle(self, *, failed: bool) -> None: """Set outcome flags from what was audible: never report completion after a failure or a dropped clause; keep suppression whenever audio was audible (no replay from the start).""" - audible = self._handle.audible - degraded = failed or self._dropped + audible, degraded = self._handle.audible, failed or self._dropped self._completed = audible and not degraded - if audible and degraded: - self._partial = True + self._partial = self._partial or (audible and degraded) self._suppress_whole_file = audible - async def _run(self) -> None: - """Drain clauses from the queue, synthesise, and write to the adapter.""" + async def _open_handle(self) -> bool: + """Open the adapter's streaming-audio handle; False when unsupported or begin failed.""" if not self.active: - return + return False if not self._adapter.supports_streaming_tts(self._chat_id, self._audio_format): logger.debug("adapter %s does not support streaming TTS", getattr(self._adapter, "name", "?")) - return + return False try: self._handle = await self._adapter.begin_streaming_tts( self._chat_id, self._audio_format, metadata=self._metadata @@ -159,27 +144,36 @@ class StreamingTTSConsumer: except Exception as exc: logger.debug("begin_streaming_tts failed: %s", exc) self._handle = None - return - if self._handle is None: + return self._handle is not None + + async def _drain(self) -> bool: + """Synthesise queued clauses until a sentinel/abort; False when a clause failed.""" + while not self._aborted: + try: + item = await asyncio.to_thread(self._queue.get, True, 0.1) + except queue.Empty: + continue + if item is _ABORT or item is _DONE or self._aborted: + break + if not isinstance(item, str): + continue + try: + await self._synthesise_and_write(item) + except Exception as exc: + logger.warning("streaming TTS clause failed: %s", exc) + self._settle(failed=True) + await self._safe_abort(str(exc)) + return False + return True + + async def _run(self) -> None: + """Drain clauses from the queue, synthesise, and write to the adapter.""" + if not await self._open_handle(): return self._suppress_whole_file = False try: - while not self._aborted: - try: - item = await asyncio.to_thread(self._queue.get, True, 0.1) - except queue.Empty: - continue - if item is _ABORT or item is _DONE or self._aborted: - break - if not isinstance(item, str): - continue - try: - await self._synthesise_and_write(item) - except Exception as exc: - logger.warning("streaming TTS clause failed: %s", exc) - self._settle(failed=True) - await self._safe_abort(str(exc)) - return + if not await self._drain(): + return if not self._aborted and self._handle is not None: try: await self._adapter.finish_streaming_tts(self._handle, interrupted=self._aborted) diff --git a/gateway/turn_lease.py b/gateway/turn_lease.py index 91b7f2ee32..3758c627b9 100644 --- a/gateway/turn_lease.py +++ b/gateway/turn_lease.py @@ -3,13 +3,12 @@ Busy guards are keyed by ROUTING KEY but the transcript is owned by SESSION_ID, and ``switch_session()`` makes key->id many-to-one (/resume from a second chat, CLI-continuity, delegation pinning, topic tip-walks): two keys ran concurrent turns on one transcript and -interleaved flushes (``user;user`` wedge). The lease serializes per RESOLVED session_id: -acquired post-resolution right before the transcript load, released in the dispatch layer's -``finally``. Release is generation-scoped and identity-checked (a stale unwind never frees a -newer turn's lease); a timed-out waiter fails CLOSED (:class:`TurnLeaseTimeoutError`, the turn -is rejected with a resend notice, never run unserialized); eviction only drops idle entries. -Known limits: CLI-continuity processes are outside this in-process lock; mid-turn compression -rotation leaves an alias window closed by :meth:`SessionTurnLeaseRegistry.rebind`. +interleaved flushes (``user;user`` wedge). The lease serializes per RESOLVED session_id: acquired +post-resolution right before the transcript load, released in the dispatch layer's ``finally``. +Release is generation-scoped and identity-checked; a timed-out waiter fails CLOSED +(:class:`TurnLeaseTimeoutError`); eviction only drops idle entries. Known limits: CLI-continuity +processes are outside this in-process lock; mid-turn compression rotation leaves an alias window +closed by :meth:`SessionTurnLeaseRegistry.rebind`. """ import asyncio @@ -58,10 +57,8 @@ class TurnLeaseToken: self.released = False def __repr__(self) -> str: # pragma: no cover - debug aid - return ( - f"TurnLeaseToken(session_id={self.session_id!r}, owner_key={self.owner_key!r}, " - f"generation={self.generation}, released={self.released})" - ) + return (f"TurnLeaseToken(session_id={self.session_id!r}, owner_key={self.owner_key!r}, " + f"generation={self.generation}, released={self.released})") class _SessionLease: @@ -100,16 +97,12 @@ class SessionTurnLeaseRegistry: return lease def _evict_idle(self) -> None: - """Drop oldest idle entries so a new lease fits under the cap; never a held/contended one. - """ - overflow = len(self._leases) - self._max_entries + 1 - if overflow <= 0: + """Drop oldest idle entries to fit a new lease under the cap; never a held/contended one.""" + if (overflow := len(self._leases) - self._max_entries + 1) <= 0: return - idle_ids = sorted( - (sid for sid, lease in self._leases.items() if lease.idle), - key=lambda sid: self._leases[sid].last_used, - ) - for sid in idle_ids[:overflow]: + idle = sorted((sid for sid, l in self._leases.items() if l.idle), + key=lambda sid: self._leases[sid].last_used) + for sid in idle[:overflow]: self._leases.pop(sid, None) async def acquire( @@ -158,15 +151,11 @@ class SessionTurnLeaseRegistry: def rebind(self, token: Optional[TurnLeaseToken], new_session_id: str) -> bool: """Alias a HELD lease onto ``new_session_id`` after mid-turn session_id rotation - (compression), so the flush target stays serialized. The SAME ``_SessionLease`` is - registered under the new id (old mapping stays until idle-evicted) — no lock state moves; - only the current holder may rebind and the token follows. If the new id already has a - live lease, log loudly and keep the old id (fail-open: a holder cannot wait mid-turn). - """ + (compression) so the flush target stays serialized: the SAME ``_SessionLease`` is registered + under the new id (old mapping stays until idle-evicted), only the current holder may rebind, + the token follows. A live lease on the new id: log loudly, keep the old id (fail-open).""" if ( - token is None - or token.released - or not new_session_id + token is None or token.released or not new_session_id or new_session_id == token.session_id ): return False