diff --git a/gateway/run_voice.py b/gateway/run_voice.py index 3afd9a272f..efa7e0c14b 100644 --- a/gateway/run_voice.py +++ b/gateway/run_voice.py @@ -1,9 +1,6 @@ -"""Voice-channel / auto-TTS methods for GatewayRunner. - -Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO. -``gateway.run`` internals are imported lazily inside method bodies (import cycle), -so ``patch("gateway.run.X")`` keeps intercepting them at call time. -""" +"""Voice-channel / auto-TTS methods for GatewayRunner (split out of ``gateway/run.py``; bound via +the MRO). ``gateway.run`` internals are imported lazily inside method bodies (import cycle), so +``patch("gateway.run.X")`` keeps intercepting them at call time.""" from __future__ import annotations @@ -18,17 +15,13 @@ import time from contextlib import suppress from difflib import SequenceMatcher from types import SimpleNamespace -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import Dict, List, Optional from gateway.config import Platform from gateway.platforms.base import MessageEvent, MessageType, build_auto_tts_output_path from gateway.session import SessionSource -if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) - from gateway.run import GatewayRunner, TurnRunner # noqa: F401 - -# Log-record parity with the origin module. -logger = logging.getLogger("gateway.run") +logger = logging.getLogger("gateway.run") # log-record parity with the origin module # 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" @@ -36,12 +29,9 @@ _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: - """``::`` 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 (else two bots in one channel + share a key and one ``/voice`` flips the other's); default keeps ``:``.""" 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}" @@ -56,8 +46,7 @@ class GatewayVoiceMixin: """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, - ) + self._handle_voice_channel_input, adapter=adapter) def _load_voice_modes(self) -> Dict[str, str]: try: @@ -70,8 +59,7 @@ class GatewayVoiceMixin: for key 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.", key, - ) + "Re-enable voice mode on that chat to rebuild the prefixed key.", key) return {k: m for k, m in items.items() if ":" in k} def _save_voice_modes(self) -> None: @@ -85,32 +73,32 @@ 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).""" + adding also clears the other set (``/voice off`` and ``/voice on``/``tts`` override).""" add_to, clear_from = (_ON_SET, _OFF_SET) if enable else (_OFF_SET, _ON_SET) - target = getattr(adapter, add_to, None) - if not isinstance(target, set): + if not isinstance(target := getattr(adapter, add_to, None), set): return if not on: target.discard(chat_id) return target.add(chat_id) - other = getattr(adapter, clear_from, None) - if isinstance(other, set): + 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.""" self._toggle_adapter_auto_tts_set(adapter, chat_id, disabled, enable=False) def _set_adapter_auto_tts_enabled(self, adapter, chat_id: str, enabled: bool) -> None: - """Update an adapter's per-chat auto-TTS opt-in set (works with ``voice.auto_tts`` off).""" self._toggle_adapter_auto_tts_set(adapter, chat_id, enabled, enable=True) + def _apply_voice_mode(self, adapter, voice_key: str, chat_id: str, mode: str) -> None: + """Record+persist ``mode``; mirror into adapter sets (``off`` -> disabled, else enabled).""" + self._voice_mode[voice_key] = mode + self._save_voice_modes() + self._toggle_adapter_auto_tts_set(adapter, chat_id, True, enable=mode != "off") + def _sync_voice_mode_state_to_adapter(self, adapter) -> None: """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``.""" + ``voice.auto_tts``; enabled (voice_only/all) / disabled (off) sets from ``_voice_mode``.""" platform = getattr(adapter, "platform", None) if not isinstance(platform, Platform): return @@ -121,9 +109,8 @@ class GatewayVoiceMixin: ] if not chat_sets: return - # Lazy import: no module-level dep from gateway -> hermes_cli. try: - from hermes_cli.config import load_config + from hermes_cli.config import load_config # lazy: no gateway -> hermes_cli module dep auto_tts_default = bool((load_config().get("voice") or {}).get("auto_tts", False)) except Exception: auto_tts_default = False @@ -132,21 +119,17 @@ class GatewayVoiceMixin: prefix = self._voice_key(platform, "", profile=getattr(adapter, "_owner_profile", None)) for chats, modes in chat_sets: chats.clear() - chats.update( - key[len(prefix):] for key, mode in self._voice_mode.items() - if mode in modes and key.startswith(prefix) - ) + chats.update(key[len(prefix):] for key, mode in self._voice_mode.items() + if mode in modes and key.startswith(prefix)) @staticmethod def _get_guild_id(event: MessageEvent) -> Optional[int]: - """Extract Discord guild_id from the raw message object.""" raw = getattr(event, "raw_message", None) if getattr(raw, "guild_id", None): # slash command interaction return int(raw.guild_id) 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.""" adapter = self._adapter_for_source(event.source) if not hasattr(adapter, "join_voice_channel"): return "Voice channels are not supported on this platform." @@ -161,14 +144,12 @@ class GatewayVoiceMixin: voice_profile = self._adapter_profile_for_source(event.source) if hasattr(adapter, "_on_voice_disconnect"): adapter._on_voice_disconnect = functools.partial( - self._handle_voice_timeout_cleanup, adapter=adapter, - ) + 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), profile=voice_profile), "off" - ) + self._voice_key(Platform.DISCORD, str(chat_id), profile=voice_profile), "off") try: success = await adapter.join_voice_channel(voice_channel) except Exception as e: @@ -184,29 +165,25 @@ class GatewayVoiceMixin: 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._set_voice_mode(self._voice_key_for_source(event.source), "all") - self._set_adapter_auto_tts_enabled(adapter, event.source.chat_id, enabled=True) - return ( - f"Joined voice channel **{voice_channel.name}**.\n" - f"I'll speak my replies and listen to you. Use /voice leave to disconnect." - ) + self._apply_voice_mode(adapter, self._voice_key_for_source(event.source), + event.source.chat_id, "all") + return (f"Joined voice channel **{voice_channel.name}**.\n" + f"I'll speak my replies and listen to you. Use /voice leave to disconnect.") async def _handle_voice_channel_leave(self, event: MessageEvent) -> str: - """Leave the Discord voice channel.""" adapter = self._adapter_for_source(event.source) guild_id = self._get_guild_id(event) - 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) - ): + 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: await adapter.leave_voice_channel(guild_id) except Exception as e: logger.warning("Error leaving voice channel: %s", e) # Always clean up state even if leave raised an exception - self._set_voice_mode(self._voice_key_for_source(event.source), "off") - self._set_adapter_auto_tts_disabled(adapter, event.source.chat_id, disabled=True) + self._apply_voice_mode(adapter, self._voice_key_for_source(event.source), + event.source.chat_id, "off") if hasattr(adapter, "_voice_input_callback"): adapter._voice_input_callback = None return "Left voice channel." @@ -216,35 +193,23 @@ class GatewayVoiceMixin: ``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) - self._set_voice_mode(self._voice_key(Platform.DISCORD, chat_id, profile=profile), "off") - self._set_adapter_auto_tts_disabled(adapter, chat_id, disabled=True) - - def _set_voice_mode(self, voice_key: str, mode: str) -> None: - """Record ``mode`` for ``voice_key`` and persist the voice-mode file.""" - self._voice_mode[voice_key] = mode - self._save_voice_modes() + key = self._voice_key(Platform.DISCORD, chat_id, + profile=getattr(adapter, "_owner_profile", None)) + self._apply_voice_mode(adapter, key, chat_id, "off") def _is_duplicate_voice_transcript(self, guild_id: int, user_id: int, transcript: str) -> bool: - """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). - """ + """Suppress repeated STT outputs for one recent utterance (voice capture can emit it twice a + few seconds apart -> a second queued run and overlapping spoken replies).""" normalized = re.sub(r"[^\w\s]", "", re.sub(r"\s+", " ", transcript).strip().lower()) if not normalized: return False - now = time.monotonic() - key = (guild_id, user_id) - recent_store = getattr(self, "_recent_voice_transcripts", None) - if not isinstance(recent_store, dict): + now, key = time.monotonic(), (guild_id, user_id) + if not isinstance(recent_store := getattr(self, "_recent_voice_transcripts", None), dict): recent_store = self._recent_voice_transcripts = {} recent = [(ts, txt) for ts, txt in recent_store.get(key, []) if now - ts <= 12.0] - if any( - prior == normalized or ( - len(prior) >= 16 and len(normalized) >= 16 - and SequenceMatcher(None, prior, normalized).ratio() >= 0.95 - ) - for _, prior in recent - ): + if any(prior == normalized or (min(len(prior), len(normalized)) >= 16 + and SequenceMatcher(None, prior, normalized).ratio() >= 0.95) + for _, prior in recent): recent_store[key] = recent return True recent_store[key] = (recent + [(now, normalized)])[-5:] @@ -261,24 +226,13 @@ class GatewayVoiceMixin: 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 + profile=getattr(adapter, "_owner_profile", None)) async def _handle_voice_channel_input( self, guild_id: int, user_id: int, transcript: str, *, adapter=None ): - """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.""" + """Handle transcribed voice from a voice channel. ``adapter`` captured the audio; under + multiplexing each profile's bot dispatches through its own adapter, never the default'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 @@ -289,10 +243,8 @@ class GatewayVoiceMixin: logger.debug("Unauthorized voice input from user %d, ignoring", user_id) return if self._is_duplicate_voice_transcript(guild_id, user_id, transcript): - logger.info( - "Suppressing duplicate voice transcript for guild=%s user=%s: %s", - guild_id, user_id, transcript[:100], - ) + logger.info("Suppressing duplicate voice transcript for guild=%s user=%s: %s", + guild_id, user_id, transcript[:100]) return # Echo the transcript into the text channel (after auth, with mention sanitization). with suppress(Exception): @@ -301,22 +253,26 @@ class GatewayVoiceMixin: safe_text = transcript[:2000].replace("@everyone", "@\u200beveryone") safe_text = safe_text.replace("@here", "@\u200bhere") await channel.send(f"**[Voice]** <@{user_id}>: {safe_text}") + # Bound text channel's channel_prompt: voice input gets the same per-channel context. + channel_prompt = None + if callable(resolver := getattr(adapter, "_resolve_channel_prompt", None)): + 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. event = MessageEvent( source=source, text=transcript, message_type=MessageType.VOICE, raw_message=SimpleNamespace(guild_id=guild_id, guild=None), - channel_prompt=self._voice_channel_prompt(adapter, text_ch_id), - ) + channel_prompt=channel_prompt) await adapter.handle_message(event) def _should_send_voice_reply( self, event: MessageEvent, response: str, agent_messages: list, already_sent: bool = False ) -> bool: """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.""" + 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 runner must do it.""" if not response or response.startswith("Error:"): return False chat_id = event.source.chat_id @@ -328,79 +284,54 @@ class GatewayVoiceMixin: adapter_auto_tts = bool(adapter._should_auto_tts_for_chat(chat_id)) # ``voice.auto_tts`` (synced into the adapter at startup) is the fallback only when the # chat has no explicit mode; the chat-level all/voice_only/off choice takes precedence. - if not ( - voice_mode == "all" - or (voice_mode == "voice_only" and is_voice_input) - or (voice_mode is None and adapter_auto_tts) - ): + if not (voice_mode == "all" or (voice_mode == "voice_only" and is_voice_input) + or (voice_mode is None and adapter_auto_tts)): logger.debug( "Auto voice reply skipped: mode=%s adapter_auto_tts=%s chat=%s platform=%s", - voice_mode, adapter_auto_tts, chat_id, event.source.platform.value, - ) + voice_mode, adapter_auto_tts, chat_id, event.source.platform.value) return False # Dedup: agent already called the TTS tool in THIS turn (from the last user message on). - 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 agent_messages[start:] if msg.get("role") == "assistant" - for tc in (msg.get("tool_calls") or []) - ): + 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 agent_messages[start:] if msg.get("role") == "assistant" + for tc in (msg.get("tool_calls") or [])): return False # Dedup: base adapter auto-TTS already handles voice input (play_tts plays in VC when - # connected), so the runner can skip — unless streaming already delivered the text - # (already_sent): then the base adapter gets None, can't run auto-TTS, and the runner must. + # connected) — unless streaming consumed the text (already_sent): then the runner must. return not (is_voice_input and not already_sent) def _should_echo_stt_transcripts(self) -> bool: - """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] = [] + """Generate TTS audio and send as a voice message before the text reply. The TTS tool + may return one combined file or several separately valid ones (combination unavailable / + over a platform limit); legacy single-file results keep working.""" + audio_path, actual_paths = None, [] try: - from tools.tts_tool import _strip_markdown_for_tts - + from tools.tts_tool import _strip_markdown_for_tts, text_to_speech_tool tts_text = _strip_markdown_for_tts(text) if not tts_text: return # 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) - actual_paths = await self._synthesize_voice_reply(tts_text, audio_path) - if actual_paths: - await self._deliver_voice_reply(event, actual_paths) + raw = await asyncio.to_thread(text_to_speech_tool, text=tts_text, + output_path=audio_path) + try: + result = json.loads(raw) + except (json.JSONDecodeError, TypeError): + logger.warning("Auto voice reply TTS returned invalid JSON: %s", + raw[:200] if raw else raw) + return + candidates = result.get("file_paths") or [result.get("file_path", audio_path)] + paths = [str(p) for p in candidates if p and os.path.isfile(p)] + if not result.get("success") or not paths: + logger.warning("Auto voice reply TTS failed: %s", result.get("error")) + return + actual_paths = 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: @@ -418,17 +349,13 @@ class GatewayVoiceMixin: for path in audio_paths: await play(guild_id, path) return - send_voice = getattr(adapter, "send_voice", None) - 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) - # so adapters that gate push notifications (Telegram "important" mode) deliver it as a - # normal notification. Clone first: the metadata is shared with typing-indicator state. + # notify=True mirrors the final-text path in platforms/base.py so notification-gating + # adapters (Telegram "important" mode) deliver it. Clone: shared w/ typing-indicator state. thread_meta = dict(self._thread_metadata_for_source(event.source, reply_anchor) or {}) thread_meta["notify"] = True for path in audio_paths: - await send_voice( - chat_id=event.source.chat_id, audio_path=path, reply_to=reply_anchor, - metadata=thread_meta, - ) + await send_voice(chat_id=event.source.chat_id, audio_path=path, reply_to=reply_anchor, + metadata=thread_meta) diff --git a/gateway/run_watchers.py b/gateway/run_watchers.py index e3492e6f5e..e5e4d5f9e9 100644 --- a/gateway/run_watchers.py +++ b/gateway/run_watchers.py @@ -1,20 +1,24 @@ -"""Session expiry / stall / catalog-refresh watcher loops for GatewayRunner. +"""Session expiry / stall / catalog-refresh watcher loops, bound onto ``GatewayRunner`` via the MRO. -Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO. -``gateway.run`` internals are imported lazily inside method bodies (import cycle), -so ``patch("gateway.run.X")`` keeps intercepting them at call time. +``gateway.run`` internals are imported lazily inside method bodies (import cycle), so +``patch("gateway.run.X")`` keeps intercepting them at call time. """ from __future__ import annotations import asyncio +import contextlib import logging import time from collections import Counter -from typing import TYPE_CHECKING, Any, Dict, Optional +from typing import Any, Dict, Optional -if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) - from gateway.run import GatewayRunner, TurnRunner # noqa: F401 +from gateway.session_stall import ( + format_session_stall_notification, + resolve_session_idle_seconds_from_activity, + should_clear_session_stall_notification, + should_emit_session_stall_notification, +) # Log-record parity with the origin module. logger = logging.getLogger("gateway.run") @@ -25,8 +29,7 @@ _SESSION_STORE_PRUNE_INTERVAL = 3600.0 # once per hour def _platform_of_key(key: str, default: str = "") -> str: """Session keys look like ``agent:main:telegram:dm:12345`` — platform is field [2].""" - parts = key.split(":") - return parts[2] if len(parts) > 2 else default + return key.split(":")[2] if key.count(":") > 1 else default async def _interruptible_sleep(runner, seconds: int) -> None: @@ -41,43 +44,32 @@ class GatewaySessionWatchersMixin: """Session expiry / stall / catalog-refresh watcher loops for GatewayRunner.""" async def _session_expiry_watcher(self, interval: int = 300): - """Finalize expired sessions (``on_session_finalize`` hooks, cached agent teardown, cache - eviction, ``expiry_finalized`` flag) and run the cache/store sweeps.""" + """Finalize expired sessions, then 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: - 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)", - len(expired), ", ".join(f"{p}:{c}" for p, c in sorted(platforms.items())), - ) + store = self.async_session_store + await store._ensure_loaded() + expired = [ + (key, entry) for key, entry in list(self.session_store._entries.items()) + if not entry.expiry_finalized and await store._is_session_expired(entry) + ] + if expired: await self._finalize_expired_sessions(expired, finalize_failures) - done = sum(1 for _, e in expired if e.expiry_finalized) - if failed := len(expired) - done: - logger.info( - "Session expiry done: %d finalized, %d pending retry", done, failed - ) - else: - logger.info("Session expiry done: %d finalized", done) await self._expiry_housekeeping() except Exception as e: logger.debug("Session expiry watcher error: %s", e) await _interruptible_sleep(self, interval) - async def _collect_expired_sessions(self) -> list: - """Return ``[(session_key, entry)]`` for expired, not-yet-finalized sessions.""" - 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.""" + platforms = Counter(_platform_of_key(k, "unknown") for k, _ in expired) + logger.info( + "Session expiry: %d sessions to finalize (%s)", + len(expired), ", ".join(f"{p}:{c}" for p, c in sorted(platforms.items())), + ) for key, entry in expired: sid = entry.session_id try: @@ -85,54 +77,47 @@ class GatewaySessionWatchersMixin: except Exception as e: count = failures[sid] = failures.get(sid, 0) + 1 if count < _MAX_FINALIZE_RETRIES: - logger.debug( - "Session finalize failed (%d/%d) for %s: %s", - count, _MAX_FINALIZE_RETRIES, sid, e, - ) + logger.debug("Session finalize failed (%d/%d) for %s: %s", + count, _MAX_FINALIZE_RETRIES, sid, e) continue logger.warning( "Session finalize gave up after %d attempts for %s: %s. " - "Marking as finalized to prevent infinite retry loop.", - count, sid, e, + "Marking as finalized to prevent infinite retry loop.", count, sid, e, ) store = self.async_session_store await store.set_expiry_finalized(entry, clear_model_override=False) 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 + done = sum(1 for _, e in expired if e.expiry_finalized) + if failed := len(expired) - done: + logger.info("Session expiry done: %d finalized, %d pending retry", done, failed) + else: + logger.info("Session expiry done: %d finalized", done) async def _finalize_expired_session(self, key: str, entry) -> None: """Run finalize hooks, tear down the cached agent, clear conversation scope, persist.""" from gateway.run import _AGENT_PENDING_SENTINEL - - try: - # Off-loop + bounded: plugin finalize hooks can block arbitrarily, and this - # watcher runs on the gateway event loop. + # Off-loop + bounded: plugin finalize hooks can block arbitrarily; this is the event loop. + with contextlib.suppress(Exception): await self._finalize_session_off_loop( - session_id=entry.session_id, - platform=_platform_of_key(key), + session_id=entry.session_id, platform=_platform_of_key(key), reason="session_expired", ) - except Exception: - pass - agent = self._agent_for_expired_session(key) + # 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, cache_lock = None, getattr(self, "_agent_cache_lock", None) # tests may lack 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 None: + state = self._peek_session_state(key) + agent = state.turn.agent if state else None 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 - # conversation-scoped dict AND boundary security state — only finalize, /new, /reset may - # do this (idle-cache eviction must NOT: a resumed turn rebuilds from those overrides). The - # persisted flag also drops the /model override — finalization is a conversation boundary. + # Evict so the AIAgent (LLM clients, tool schemas, memory refs) can be GC'd, then drop + # every conversation-scoped dict AND boundary security state — only finalize, /new, /reset + # may (idle-cache eviction must NOT: a resumed turn rebuilds from those overrides). The + # persisted flag also drops the /model override: finalization is a conversation boundary. self._evict_cached_agent(key) self._clear_conversation_scope(key, reason="expiry_finalized") await self.async_session_store.set_expiry_finalized(entry) @@ -140,8 +125,7 @@ class GatewaySessionWatchersMixin: async def _expiry_housekeeping(self) -> None: """Idle/pressure agent-cache sweeps plus the hourly SessionStore prune.""" - # Sweep agents idle beyond the TTL regardless of session reset policy: long / "never" - # reset windows would otherwise pin memory for the gateway's life. + # Sweep idle agents regardless of reset policy: long/"never" windows would pin memory. try: if evicted := self._sweep_idle_cached_agents(): logger.info("Agent cache idle sweep: evicted %d agent(s)", evicted) @@ -152,16 +136,13 @@ class GatewaySessionWatchersMixin: self._sweep_agent_cache_under_pressure() except Exception as e: logger.debug("Agent cache pressure sweep failed: %s", e) - # Prune stale SessionStore entries; the in-memory dict (and sessions.json) would - # otherwise grow unbounded with many rotating chats / threads / users. - last_prune = getattr(self, "_last_session_store_prune_ts", 0.0) - if time.time() - last_prune > _SESSION_STORE_PRUNE_INTERVAL: + # Prune stale SessionStore entries: the dict + sessions.json otherwise grow unbounded. + prune_ts = getattr(self, "_last_session_store_prune_ts", 0.0) # tests may omit + if time.time() - prune_ts > _SESSION_STORE_PRUNE_INTERVAL: try: max_age = int(getattr(self.config, "session_store_max_age_days", 0) or 0) - if max_age > 0: - pruned = await self.async_session_store.prune_old_entries(max_age) - if pruned: - logger.info("SessionStore prune: dropped %d stale entries", pruned) + if max_age > 0 and (n := await self.async_session_store.prune_old_entries(max_age)): + logger.info("SessionStore prune: dropped %d stale entries", n) except Exception as e: logger.debug("SessionStore prune failed: %s", e) self._last_session_store_prune_ts = time.time() @@ -171,19 +152,8 @@ class GatewaySessionWatchersMixin: from gateway.run import _float_env return _float_env("HERMES_SESSION_STALL_TIMEOUT", 300) - 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()) - for amap in maps: - for adapter in list(amap.values()): - if adapter is not None and id(adapter) not in seen: - seen.add(id(adapter)) - yield adapter - def _session_activity_for_stall(self, session_key: str) -> Optional[dict]: - """Activity snapshot for stall progress: the single source is - ``AIAgent.get_activity_summary()``; no turn-start or pending-inbound clocks.""" + """Stall-progress snapshot from ``AIAgent.get_activity_summary()`` only; no other 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: @@ -195,55 +165,41 @@ class GatewaySessionWatchersMixin: return summary if isinstance(summary, dict) else None def _stall_candidates(self) -> Dict[str, tuple[Any, Any]]: - """Map session_key -> (adapter, pending event) from adapter pending slots, then from the - runner's overflow queues (first occurrence wins).""" + """session_key -> (adapter, pending event) from every live adapter's pending slot (default + + multiplex profiles, deduped by identity), then the overflow queues; first one wins.""" candidates: Dict[str, tuple[Any, Any]] = {} - for adapter in self._iter_gateway_adapters(): - pending_slot = getattr(adapter, "_pending_messages", None) or {} - for session_key, event in list(pending_slot.items()): + maps = (getattr(self, "adapters", {}), *getattr(self, "_profile_adapters", {}).values()) + adapters = {id(a): a for m in maps for a in list(m.values()) if a is not None} + for adapter in adapters.values(): + pending = getattr(adapter, "_pending_messages", None) or {} + for session_key, event in list(pending.items()): if session_key and session_key not in candidates and event is not None: candidates[session_key] = (adapter, event) for session_key, overflow in list((getattr(self, "_queued_events", None) or {}).items()): if not session_key or session_key in candidates or not overflow: continue source = getattr(overflow[0], "source", None) - adapter = self._adapter_for_source(source) if source is not None else None - if adapter is not None: + if source is not None and (adapter := self._adapter_for_source(source)) is not None: candidates[session_key] = (adapter, overflow[0]) return candidates - def _session_still_pending(self, adapter, session_key: str) -> bool: - return ( - (getattr(adapter, "_pending_messages", None) or {}).get(session_key) is not None - or bool((getattr(self, "_queued_events", None) or {}).get(session_key)) - ) - 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).""" - from gateway.session_stall import ( - resolve_session_idle_seconds_from_activity, - should_clear_session_stall_notification, - should_emit_session_stall_notification, - ) - - notified_map = getattr(self, "_session_stall_notified", None) - if notified_map is None: - notified_map = self._session_stall_notified = {} + """Notify once per stall episode for pending inbound sessions; returns notices sent.""" + if getattr(self, "_session_stall_notified", None) is None: # tests may build bare runners + self._session_stall_notified = {} + notified_map = self._session_stall_notified 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) idle_seconds = resolve_session_idle_seconds_from_activity(activity, now=now) - already = bool(notified_map.get(session_key)) if should_clear_session_stall_notification( timeout_seconds=timeout_seconds, idle_seconds=idle_seconds, has_pending_inbound=True ): notified_map.pop(session_key, None) - already = False if idle_seconds is None or not should_emit_session_stall_notification( timeout_seconds=timeout_seconds, idle_seconds=idle_seconds, - has_pending_inbound=True, already_notified=already, + has_pending_inbound=True, already_notified=bool(notified_map.get(session_key)), ): continue if await self._notify_session_stall( @@ -252,51 +208,16 @@ class GatewaySessionWatchersMixin: ): sent += 1 # Drop latches for sessions that no longer appear in any pending map. - for key in list(notified_map.keys()): - if key not in candidates: - notified_map.pop(key, None) + for key in [k for k in notified_map if k not in candidates]: + notified_map.pop(key, None) return sent - 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 - - try: - metadata = self._thread_metadata_for_source(source) - notice = format_session_stall_notification(idle_seconds) - # Bound the send: a wedged adapter transport (network hang, dead websocket) must not - # block the watcher pass — siblings would go unevaluated and the watcher stop. - result = await asyncio.wait_for( - adapter.send(str(source.chat_id), notice, metadata=metadata), - timeout=_STALL_NOTIFY_SEND_TIMEOUT_SECONDS, - ) - except asyncio.TimeoutError: - logger.warning( - "Session stall notify send timed out after %.0fs for %s; will retry next tick", - _STALL_NOTIFY_SEND_TIMEOUT_SECONDS, session_key, - ) - return False - except Exception as exc: - logger.warning("Session stall notify failed for %s: %s", session_key, exc) - return False - # Adapters often return SendResult(success=False) instead of raising. - if result is not None and getattr(result, "success", True) is False: - logger.warning( - "Session stall notify failed for %s: %s", - session_key, getattr(result, "error", "send returned success=False"), - ) - 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: + 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 - + from gateway.run import _STALL_NOTIFY_SEND_TIMEOUT_SECONDS 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)", @@ -313,19 +234,38 @@ class GatewaySessionWatchersMixin: # 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) + still_pending = ( + (getattr(adapter, "_pending_messages", None) or {}).get(session_key) is not None + or bool((getattr(self, "_queued_events", None) or {}).get(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, - ) + 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): + try: + metadata = self._thread_metadata_for_source(source) + notice = format_session_stall_notification(idle_seconds) + # Bound the send: a wedged adapter transport (network hang, dead websocket) must not + # block the watcher pass — siblings would go unevaluated and the watcher stop. + result = await asyncio.wait_for( + adapter.send(str(source.chat_id), notice, metadata=metadata), + timeout=_STALL_NOTIFY_SEND_TIMEOUT_SECONDS, + ) + # Adapters often return SendResult(success=False) instead of raising. + if result is not None and getattr(result, "success", True) is False: + raise RuntimeError(getattr(result, "error", "send returned success=False")) + except asyncio.TimeoutError: + logger.warning( + "Session stall notify send timed out after %.0fs for %s; will retry next tick", + _STALL_NOTIFY_SEND_TIMEOUT_SECONDS, session_key, + ) + return False + except Exception as exc: + logger.warning("Session stall notify failed for %s: %s", session_key, exc) return False notified_map[session_key] = True return True @@ -334,7 +274,6 @@ class GatewaySessionWatchersMixin: """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.""" from hermes_cli.model_catalog import refresh_catalogs, refresh_interval_seconds - await asyncio.sleep(30) # let startup settle while self._running: try: @@ -350,10 +289,9 @@ class GatewaySessionWatchersMixin: await asyncio.sleep(min(30.0, max(0.0, deadline - time.monotonic()))) async def _session_stall_watcher(self, interval: float = 30.0): - """Periodic pending-inbound + stale-activity stall watchdog. Progress comes only from + """Pending-inbound + stale-activity stall watchdog. Progress comes only from ``get_activity_summary()``; pending inbound is a notify policy gate, not a progress clock. - Notify-only: does not kill the turn (contrast ``gateway_timeout`` / ``shutdown_watchdog``). - """ + Notify-only: never kills the turn (contrast ``gateway_timeout`` / ``shutdown_watchdog``).""" # Short initial delay so startup reconnect noise does not false-fire. await asyncio.sleep(min(30.0, max(1.0, float(interval)))) while self._running: diff --git a/gateway/sticker_cache.py b/gateway/sticker_cache.py index 40b79d7a18..2760ef030c 100644 --- a/gateway/sticker_cache.py +++ b/gateway/sticker_cache.py @@ -4,15 +4,12 @@ Stickers are described via the vision tool once and cached by file_unique_id (``~/.hermes/sticker_cache.json``) so the same image is never re-analyzed. """ -import contextlib import json -import os -import tempfile import time from typing import Optional from hermes_cli.config import get_hermes_home - +from utils import atomic_json_write CACHE_PATH = get_hermes_home() / "sticker_cache.json" @@ -31,19 +28,7 @@ def _load_cache() -> dict: def _save_cache(cache: dict) -> None: - """Write the cache atomically (temp file + fsync + replace).""" - CACHE_PATH.parent.mkdir(parents=True, exist_ok=True) - fd, tmp_path = tempfile.mkstemp(dir=str(CACHE_PATH.parent), suffix=".tmp") - try: - with os.fdopen(fd, "w", encoding="utf-8") as f: - json.dump(cache, f, indent=2, ensure_ascii=False) - f.flush() - os.fsync(f.fileno()) - os.replace(tmp_path, str(CACHE_PATH)) - except BaseException: - with contextlib.suppress(OSError): - os.unlink(tmp_path) - raise + atomic_json_write(CACHE_PATH, cache) def get_cached_description(file_unique_id: str) -> Optional[dict]: @@ -51,29 +36,28 @@ def get_cached_description(file_unique_id: str) -> Optional[dict]: return _load_cache().get(file_unique_id) -def cache_sticker_description(file_unique_id: str, description: str, emoji: str = "", set_name: str = "") -> None: +def cache_sticker_description( + file_unique_id: str, description: str, emoji: str = "", set_name: str = "" +) -> None: """Store a vision-generated description under Telegram's stable sticker id.""" - cache = _load_cache() - cache[file_unique_id] = {"description": description, "emoji": emoji, "set_name": set_name, "cached_at": time.time()} - _save_cache(cache) + entry = {"description": description, "emoji": emoji, "set_name": set_name, + "cached_at": time.time()} + _save_cache({**_load_cache(), file_unique_id: entry}) def build_sticker_injection(description: str, emoji: str = "", set_name: str = "") -> str: """Warm-style injection text, e.g. ``[The user sent a sticker 😀 from "MyPack"~ It shows: "A cat waving" (=^.w.^=)]``. - ``set_name`` is only shown together with an emoji. - """ + ``set_name`` is only shown together with an emoji.""" context = f" {emoji}" if emoji else "" if set_name and emoji: - context += f" from \"{set_name}\"" - return f"[The user sent a sticker{context}~ It shows: \"{description}\" (=^.w.^=)]" + context += f' from "{set_name}"' + return f'[The user sent a sticker{context}~ It shows: "{description}" (=^.w.^=)]' def build_animated_sticker_injection(emoji: str = "") -> str: """Injection text for animated/video stickers we can't analyze.""" if emoji: - return ( - f"[The user sent an animated sticker {emoji}~ " - f"I can't see animated ones yet, but the emoji suggests: {emoji}]" - ) + return (f"[The user sent an animated sticker {emoji}~ " + f"I can't see animated ones yet, but the emoji suggests: {emoji}]") return "[The user sent an animated sticker~ I can't see animated ones yet]" diff --git a/gateway/streaming_tts_consumer.py b/gateway/streaming_tts_consumer.py index a23b68dcfa..3528b8f250 100644 --- a/gateway/streaming_tts_consumer.py +++ b/gateway/streaming_tts_consumer.py @@ -27,37 +27,25 @@ _DONE = object() class StreamingTTSConsumer: """Consumes LLM text deltas and produces streaming PCM audio for an adapter.""" - def __init__( - self, - adapter: Any, - chat_id: str, - tts_config: Dict[str, Any], - loop: asyncio.AbstractEventLoop, - *, - metadata: Optional[Dict[str, Any]] = None, - audio_format: Optional[AudioFormat] = None, - ) -> None: + def __init__(self, adapter: Any, chat_id: str, tts_config: Dict[str, Any], + loop: asyncio.AbstractEventLoop, *, metadata: Optional[Dict[str, Any]] = None, + audio_format: Optional[AudioFormat] = None) -> None: from tools.tts_streaming import SentenceChunker, resolve_streaming_provider - 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() - if self._streamer is not None: - self._audio_format = AudioFormat(**{ - f: int(getattr(self._streamer, f, getattr(AudioFormat, f))) - for f in ("sample_rate", "channels", "sample_width") - }) - else: - self._audio_format = audio_format or AudioFormat() + self._audio_format = audio_format or AudioFormat() if self._streamer is None else ( + AudioFormat(**{f: int(getattr(self._streamer, f, getattr(AudioFormat, f))) + for f in ("sample_rate", "channels", "sample_width")}) + ) # Thread-safe queue of completed clauses plus the _DONE/_ABORT sentinels. self._queue: "queue.Queue[Any]" = queue.Queue(maxsize=256) self._handle: Optional[StreamingTTSHandle] = None + self._task: Optional[asyncio.Task] = None # drain task, created once by start() self._completed = self._partial = self._aborted = False self._finished = self._dropped = self._suppress_whole_file = False - self._task: Optional[asyncio.Task] = None - self._lock = threading.Lock() - self._strip_markdown = None # lazily imported to avoid import cycles + self._lock, self._strip_markdown = threading.Lock(), None # stripper lazily imported active = property(lambda self: self._streamer is not None) # usable streaming provider completed = property(lambda self: self._completed) # streaming audio fully delivered @@ -82,21 +70,19 @@ class StreamingTTSConsumer: """Receive a text delta from the agent. Non-blocking.""" if self._aborted or not self.active or self._finished: return - self._enqueue_clauses( - self._chunker.feed(text), "streaming TTS queue full, dropping clause", log_errors=True - ) + self._enqueue_clauses(self._chunker.feed(text), "streaming TTS queue full, dropping clause", + log_errors=True) def finish(self) -> None: - """Signal end-of-text, flush the chunker tail, then enqueue ``_DONE`` after all flushed + """Signal end-of-text: flush the chunker tail, then enqueue ``_DONE`` after all flushed clauses so the drain loop ends deterministically without racing a late ``on_delta``.""" if self._finished: return self._finished = True if self._aborted or not self.active: return - self._enqueue_clauses( - self._chunker.flush(), "streaming TTS queue full while flushing tail", log_errors=False - ) + self._enqueue_clauses(self._chunker.flush(), "streaming TTS queue full while flushing tail", + log_errors=False) # The load-bearing _DONE sentinel must never be lost: evict clauses until it fits. while not self._put_sentinel(_DONE, mark_dropped=True): pass @@ -104,11 +90,9 @@ class StreamingTTSConsumer: def _put_sentinel(self, sentinel, *, mark_dropped: bool) -> bool: """Try to enqueue a sentinel, evicting one queued item when the queue is full. Returns True when the caller should stop retrying: enqueued, or (abort path) nothing left to evict.""" - try: + with contextlib.suppress(queue.Full): self._queue.put_nowait(sentinel) return True - except queue.Full: - pass try: self._queue.get_nowait() except queue.Empty: @@ -135,9 +119,8 @@ class StreamingTTSConsumer: if not self.active: 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", "?"), - ) + name = getattr(self._adapter, "name", "?") + logger.debug("adapter %s does not support streaming TTS", name) return False try: self._handle = await self._adapter.begin_streaming_tts( @@ -148,39 +131,33 @@ class StreamingTTSConsumer: self._handle = 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.""" + """Drain clauses until a sentinel/abort, synthesise + write each, then finalise the stream; + a clause or finalise failure settles the outcome flags and aborts the adapter stream.""" if not await self._open_handle(): return self._suppress_whole_file = False try: - if not await self._drain(): - return + 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 self._aborted and self._handle is not None: try: - await self._adapter.finish_streaming_tts( - self._handle, interrupted=self._aborted, - ) + handle, interrupted = self._handle, self._aborted + await self._adapter.finish_streaming_tts(handle, interrupted=interrupted) except Exception as exc: logger.debug("finish_streaming_tts error: %s", exc) self._settle(failed=True) @@ -199,8 +176,13 @@ class StreamingTTSConsumer: """Synthesise one clause via the streamer and write PCM chunks.""" if self._handle is None or self._handle.aborted or self._streamer is None: return - cleaned = self._strip_markdown_for_tts(clause) - if not cleaned.strip(): + if self._strip_markdown is None: # lazy import: tools.tts_tool would cycle at module load + try: + from tools.tts_tool import _strip_markdown_for_tts as _strip + self._strip_markdown = _strip + except ImportError: + self._strip_markdown = lambda t: t # noqa: E731 + if not (cleaned := self._strip_markdown(clause).strip()): return iterator = iter(self._streamer.stream(cleaned)) while True: @@ -213,18 +195,7 @@ class StreamingTTSConsumer: was_audible = self._handle.audible await self._adapter.write_streaming_tts(self._handle, chunk) if not was_audible: - self._handle.audible = True - self._suppress_whole_file = True - - def _strip_markdown_for_tts(self, text: str) -> str: - """Lazy-import and apply the TTS markdown stripper.""" - if self._strip_markdown is None: - try: - from tools.tts_tool import _strip_markdown_for_tts as _strip - self._strip_markdown = _strip - except ImportError: - self._strip_markdown = lambda t: t # noqa: E731 - return self._strip_markdown(text).strip() + self._handle.audible = self._suppress_whole_file = True async def _safe_abort(self, reason: str) -> None: """Abort the adapter stream, swallowing errors (idempotent).""" @@ -244,10 +215,7 @@ class StreamingTTSConsumer: return self._aborted = True # The load-bearing _ABORT sentinel must reach the queue even when full: evict to make room. - for _attempt in range(3): - if self._put_sentinel(_ABORT, mark_dropped=False): - break - else: + if not any(self._put_sentinel(_ABORT, mark_dropped=False) for _ in range(3)): logger.debug("streaming TTS _ABORT sentinel could not be enqueued") if self._handle is not None and not self._handle.aborted: with contextlib.suppress(Exception): diff --git a/gateway/turn_context.py b/gateway/turn_context.py index 477c8e8b04..4a0f6cc84e 100644 --- a/gateway/turn_context.py +++ b/gateway/turn_context.py @@ -1,10 +1,8 @@ """Per-turn context shared between ``GatewayRunner._run_agent_inner`` and ``TurnRunner``. -Extraction seam for what used to be closures over ~20 locals: each closed-over local is a -field. Invariants: fields are written once by ``_run_agent_inner`` while wiring the turn; -``message`` (formerly ``nonlocal``) is the only rebindable field; other mutable state keeps -single-element-list containers so mutation stays visible to the outer body; -``_run_still_current`` stays a callable capturing ``self``/``session_key``/``run_generation``. +Each ex-closure local is a field, written once by ``_run_agent_inner`` while wiring the turn. +``message`` (ex-``nonlocal``) is the only rebindable field; other mutable state uses +single-element lists so mutation stays visible to the outer body. """ from __future__ import annotations @@ -15,8 +13,6 @@ from typing import Any, Callable, List, Optional @dataclass class TurnContext: - """Closed-over locals of ``_run_agent_inner`` needed by ``TurnRunner``.""" - # read-only turn identity / wiring source: Any = None _run_still_current: Callable[[], bool] = None # type: ignore[assignment] @@ -35,15 +31,12 @@ class TurnContext: repeat_count: list = field(default_factory=lambda: [0]) long_tool_hint_fired: list = field(default_factory=lambda: [False]) agent_holder: list = field(default_factory=lambda: [None]) - # constants / cleanup bookkeeping _LONG_TOOL_THRESHOLD_S: float = 30.0 _cleanup_progress: bool = False _cleanup_msg_ids: List[str] = field(default_factory=list) - # progress threading metadata (assigned before send_progress_messages runs) _progress_metadata: Optional[dict] = None _progress_reply_to: Optional[Any] = None - # run_sync seam: the ex-``nonlocal`` turn message (rebindable) - message: Optional[str] = None + message: Optional[str] = None # the only rebindable field # turn parameters / config snapshots (read-only in run_sync) history: Any = None context_prompt: Optional[str] = None @@ -55,14 +48,12 @@ class TurnContext: process_baseline: frozenset[str] = field(default_factory=frozenset) _interrupt_depth: int = 0 event_message_id: Optional[str] = None - # Raw platform id of the INBOUND user message (event.message_id), distinct from the - # event_message_id reply/thread anchor; stamped as platform_message_id on the user turn. + # Raw inbound platform id (not the event_message_id reply anchor); stamped on the user turn. inbound_message_id: Optional[str] = None moa_config: Optional[dict] = None persist_user_message: Optional[Any] = None persist_user_timestamp: Optional[float] = None - # display_kind for the persisted user row of a self-injected turn (MessageEvent.internal), - # e.g. "internal_notification". DB-only presentation metadata; never sent to the provider. + # display_kind of the persisted user row for a self-injected turn; DB-only, never sent. persist_user_display_kind: Optional[str] = None user_config: Any = None enabled_toolsets: Any = None @@ -70,10 +61,8 @@ class TurnContext: log_mode_enabled: bool = False interim_assistant_messages_enabled: bool = False needs_progress_queue: bool = False - # lazy-imported callables captured from the outer body AIAgent: Any = None resolve_display_setting: Any = None - # mutable holder cells (shared-list pattern) result_holder: list = field(default_factory=lambda: [None]) tools_holder: list = field(default_factory=lambda: [None]) stream_consumer_holder: list = field(default_factory=lambda: [None]) @@ -82,20 +71,19 @@ class TurnContext: _voice_ack_fired: list = field(default_factory=lambda: [False]) _voice_ack_guild: list = field(default_factory=lambda: [None]) _voice_ack_loop: Any = None - # hook / status bridge wiring (published at original binding sites) + # hook / status bridge wiring _loop_for_step: Any = None _hooks_ref: Any = None _status_adapter: Any = None _status_chat_id: Any = None _status_thread_metadata: Optional[dict] = None - # extracted sibling callbacks (bound TurnRunner methods read via ctx) + # bound TurnRunner callbacks read via ctx progress_callback: Optional[Callable] = None voice_ack_callback: Optional[Callable] = None _step_callback_sync: Optional[Callable] = None _event_callback_sync: Optional[Callable] = None _status_callback_sync: Optional[Callable] = None - # Slack-native task-card progress (opt-in via adapter.native_task_cards_enabled()); - # ID-bearing callbacks so tool starts/completions correlate by tool-call ID + # Slack-native task cards (opt-in); ID-bearing callbacks correlate start/complete by call ID _native_slack_task_cards: bool = False native_tool_start_callback: Optional[Callable] = None native_tool_complete_callback: Optional[Callable] = None diff --git a/gateway/turn_lease.py b/gateway/turn_lease.py index 3758c627b9..1f65614e66 100644 --- a/gateway/turn_lease.py +++ b/gateway/turn_lease.py @@ -2,13 +2,11 @@ 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 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`. +delegation pinning, topic tip-walks), so two keys could interleave flushes on one transcript +(``user;user`` wedge). The lease serializes per RESOLVED session_id: acquired right before the +transcript load, released in the dispatch layer's ``finally``; identity-checked release; a +timed-out waiter fails CLOSED (:class:`TurnLeaseTimeoutError`); only idle entries evict. Limits: +CLI-continuity processes are outside this lock; mid-turn rotation alias is closed by ``rebind``. """ import asyncio @@ -32,18 +30,13 @@ def _holder_desc(holder: Optional["TurnLeaseToken"]) -> tuple: class TurnLeaseTimeoutError(TimeoutError): - """The lease stayed held for the caller's full wait budget (fail-closed: the - caller must not enter the transcript load/run/flush region).""" + """Lease held for the full wait budget; fail-closed: caller must not enter the turn region.""" - def __init__( - self, session_id: str, *, owner_key: str, generation: int, wait_seconds: float - ) -> None: + def __init__(self, session_id: str, *, owner_key: str, generation: int, wait_seconds: float): self.session_id, self.owner_key = session_id, owner_key self.generation, self.wait_seconds = generation, wait_seconds - super().__init__( - f"turn lease wait timed out after {wait_seconds:.0f}s on session " - f"{session_id} for routing key {owner_key} (gen {generation})" - ) + super().__init__(f"turn lease wait timed out after {wait_seconds:.0f}s on session " + f"{session_id} for routing key {owner_key} (gen {generation})") class TurnLeaseToken: @@ -67,9 +60,7 @@ class _SessionLease: def __init__(self) -> None: self.lock = asyncio.Lock() self.holder: Optional[TurnLeaseToken] = None - self.acquired_at = 0.0 - self.last_used = time.time() - self.pending_acquires = 0 + self.acquired_at, self.last_used, self.pending_acquires = 0.0, time.time(), 0 @property def idle(self) -> bool: @@ -85,12 +76,8 @@ class SessionTurnLeaseRegistry: self._leases: Dict[str, _SessionLease] = {} self._max_entries = max(1, int(max_entries)) - def __len__(self) -> int: - return len(self._leases) - def _get_or_create(self, session_id: str) -> _SessionLease: - lease = self._leases.get(session_id) - if lease is None: + if (lease := self._leases.get(session_id)) is None: self._evict_idle() lease = self._leases[session_id] = _SessionLease() lease.last_used = time.time() @@ -122,8 +109,7 @@ class SessionTurnLeaseRegistry: "are mapped to one session_id (#64934); serializing this turn behind the previous " "turn's flush", session_id, owner_key, generation, *_holder_desc(lease.holder), - time.time() - lease.acquired_at if lease.acquired_at else -1.0, - ) + time.time() - lease.acquired_at if lease.acquired_at else -1.0) # Lock.release() wakes a waiter while leaving the lock momentarily unlocked. Count every # in-progress acquire across that handoff (even apparently-uncontended ones — wait_for() # may schedule them before the lock coroutine runs) so eviction cannot orphan the old @@ -136,11 +122,9 @@ class SessionTurnLeaseRegistry: "turn lease wait timed out after %.0fs on session %s (waiter: routing key %s gen " "%s; holder: routing key %s gen %s) — failing closed: refusing to run this turn " "UNSERIALIZED against the still-held lease", - wait, session_id, owner_key, generation, *_holder_desc(lease.holder), - ) + wait, session_id, owner_key, generation, *_holder_desc(lease.holder)) raise TurnLeaseTimeoutError( - session_id, owner_key=owner_key, generation=generation, wait_seconds=wait - ) from None + session_id, owner_key=owner_key, generation=generation, wait_seconds=wait) from None finally: lease.pending_acquires -= 1 # Lock held and no await before holder publication, so the lease cannot become @@ -150,17 +134,14 @@ class SessionTurnLeaseRegistry: return token 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), 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 - or new_session_id == token.session_id - ): + """Alias a HELD lease onto ``new_session_id`` after mid-turn rotation (compression) so the + flush target stays serialized: the SAME ``_SessionLease`` is registered under the new id + (old mapping idle-evicts later), only the 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 + or new_session_id == token.session_id): return False - lease = self._leases.get(token.session_id) - if lease is None or lease.holder is not token: + if (lease := self._leases.get(token.session_id)) is None or lease.holder is not token: return False existing = self._leases.get(new_session_id) if existing is not None and existing is not lease and not existing.idle: @@ -170,8 +151,7 @@ class SessionTurnLeaseRegistry: "gen %s) — keeping the lease on the old id; transcript writes on %s may " "interleave (#64934 rotation-alias edge)", token.session_id, new_session_id, token.owner_key, token.generation, - *_holder_desc(existing.holder), new_session_id, - ) + *_holder_desc(existing.holder), new_session_id) return False self._leases[new_session_id] = lease lease.last_used = time.time() @@ -184,15 +164,11 @@ class SessionTurnLeaseRegistry: if token is None or token.released: return False token.released = True - lease = self._leases.get(token.session_id) - if lease is None: + if (lease := self._leases.get(token.session_id)) is None: return False if lease.holder is not token: - logger.debug( - "turn lease release skipped on session %s: token (key %s gen %s) is not the " - "current holder", - token.session_id, token.owner_key, token.generation, - ) + logger.debug("turn lease release skipped on session %s: token (key %s gen %s) is not " + "the current holder", token.session_id, token.owner_key, token.generation) return False lease.holder, lease.acquired_at, lease.last_used = None, 0.0, time.time() if lease.lock.locked():