diff --git a/gateway/run_voice.py b/gateway/run_voice.py index 2157c92912..aafbaedb55 100644 --- a/gateway/run_voice.py +++ b/gateway/run_voice.py @@ -7,20 +7,20 @@ so ``patch("gateway.run.X")`` keeps intercepting them at call time. from __future__ import annotations -import logging -from typing import TYPE_CHECKING import asyncio import functools import json +import logging import os import re import sys import time from contextlib import suppress +from typing import TYPE_CHECKING, 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 -from typing import Any, Awaitable, Callable, Dict, List, Optional, cast if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) from gateway.run import GatewayRunner, TurnRunner # noqa: F401 @@ -28,161 +28,123 @@ if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) # Log-record parity with the origin module. 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" + class GatewayVoiceMixin: """Voice-channel / auto-TTS methods for GatewayRunner.""" - def _voice_key( - self, platform: Platform, chat_id: str, profile: Optional[str] = None - ) -> str: - """Return a platform-namespaced key for voice mode state. - - Under multiplexing the key is ``::`` (profile whose bot speaks); - the default profile keeps ``:`` so persisted state stays valid. Otherwise - two bots in one Discord channel share a key and one profile's ``/voice`` flips the other's. + 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. """ base = f"{platform.value}:{chat_id}" profile = profile.strip() if isinstance(profile, str) else "" - if not profile or profile == "default": - return base - return f"{profile}:{base}" + 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, namespaced by its transport owner. - - Voice mode belongs to the (bot, chat) pair, so the namespace is the profile that OWNS the - receiving adapter (matching ``_sync_voice_mode_state_to_adapter``), not the routed profile. - """ - return self._voice_key( - source.platform, - source.chat_id, - profile=self._adapter_profile_for_source(source), - ) + """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.""" + 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"): - adapter._voice_input_callback = functools.partial( - self._handle_voice_channel_input, adapter=adapter - ) + cb = functools.partial(self._handle_voice_channel_input, adapter=adapter) + adapter._voice_input_callback = cb 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 {} - - if not isinstance(data, dict): - return {} - - valid_modes = {"off", "voice_only", "all"} result = {} - for chat_id, mode in data.items(): - if mode not in valid_modes: + for chat_id, mode in (data.items() if isinstance(data, dict) else ()): + if mode not in {"off", "voice_only", "all"}: continue - key = str(chat_id) - # Skip legacy unprefixed keys (warn and skip) - if ":" not in key: + 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.", - key, + "Re-enable voice mode on that chat to rebuild the prefixed key.", str(chat_id), ) - continue - result[key] = mode return result def _save_voice_modes(self) -> None: try: self._VOICE_MODE_PATH.parent.mkdir(parents=True, exist_ok=True) - self._VOICE_MODE_PATH.write_text( - json.dumps(self._voice_mode, indent=2), encoding="utf-8" - ) + payload = json.dumps(self._voice_mode, indent=2) + self._VOICE_MODE_PATH.write_text(payload, encoding="utf-8") except OSError as e: logger.warning("Failed to save voice modes: %s", e) @staticmethod - def _toggle_adapter_auto_tts_set(adapter, chat_id: str, on: bool, *, add_to: str, clear_from: str) -> None: - """Add/discard ``chat_id`` in the adapter's ``add_to`` set; adding also clears it from ``clear_from``. - - ``/voice off`` and an explicit ``/voice on``/``/voice tts`` are hard overrides of each other.""" + 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_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 on: + if not on: + target.discard(chat_id) + else: target.add(chat_id) - other = getattr(adapter, clear_from, None) if isinstance(other, set): other.discard(chat_id) - else: - target.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, add_to="_auto_tts_disabled_chats", clear_from="_auto_tts_enabled_chats" - ) + 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 (auto-TTS even when ``voice.auto_tts`` is False).""" - self._toggle_adapter_auto_tts_set( - adapter, chat_id, enabled, add_to="_auto_tts_enabled_chats", clear_from="_auto_tts_disabled_chats" - ) + """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 _sync_voice_mode_state_to_adapter(self, adapter) -> None: - """Restore persisted /voice state into a live platform adapter. - - Sets ``_auto_tts_default`` (from ``voice.auto_tts``) and, from ``self._voice_mode``, - ``_auto_tts_enabled_chats`` (modes ``voice_only``/``all``) and ``_auto_tts_disabled_chats`` - (mode ``off``). + """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``. """ platform = getattr(adapter, "platform", None) if not isinstance(platform, Platform): return - - disabled_chats = getattr(adapter, "_auto_tts_disabled_chats", None) - enabled_chats = getattr(adapter, "_auto_tts_enabled_chats", None) - if not isinstance(disabled_chats, set) and not isinstance(enabled_chats, set): + chat_sets = [ + (getattr(adapter, name, None), modes) + for name, modes in ((_OFF_SET, {"off"}), (_ON_SET, {"voice_only", "all"})) + ] + chat_sets = [(chats, modes) for chats, modes in chat_sets if isinstance(chats, set)] + if not chat_sets: return - - # Push the global voice.auto_tts default (config.yaml) onto the adapter. - # Lazy import to avoid adding a module-level dep from gateway → hermes_cli. + # Lazy import: no module-level dep from gateway -> hermes_cli. try: - from hermes_cli.config import load_config as _load_full_config - _full_cfg = _load_full_config() - _auto_tts_default = bool( - (_full_cfg.get("voice") or {}).get("auto_tts", False) - ) + from hermes_cli.config import load_config + auto_tts_default = bool((load_config().get("voice") or {}).get("auto_tts", False)) except Exception: - _auto_tts_default = False + auto_tts_default = False if hasattr(adapter, "_auto_tts_default"): - adapter._auto_tts_default = _auto_tts_default - + adapter._auto_tts_default = auto_tts_default prefix = self._voice_key(platform, "", profile=getattr(adapter, "_owner_profile", None)) - if isinstance(disabled_chats, set): - disabled_chats.clear() - disabled_chats.update( + for chats, modes in chat_sets: + chats.clear() + chats.update( key[len(prefix):] for key, mode in self._voice_mode.items() - if mode == "off" and key.startswith(prefix) - ) - if isinstance(enabled_chats, set): - enabled_chats.clear() - enabled_chats.update( - key[len(prefix):] for key, mode in self._voice_mode.items() - if mode in {"voice_only", "all"} and key.startswith(prefix) + 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 raw is None: - return None - # Slash command interaction - if hasattr(raw, "guild_id") and raw.guild_id: + if getattr(raw, "guild_id", None): # slash command interaction return int(raw.guild_id) - # Regular message - if hasattr(raw, "guild") and raw.guild: + if getattr(raw, "guild", None): # regular message return raw.guild.id return None @@ -191,72 +153,62 @@ class GatewayVoiceMixin: adapter = self._adapter_for_source(event.source) if not hasattr(adapter, "join_voice_channel"): return "Voice channels are not supported on this platform." - guild_id = self._get_guild_id(event) if not guild_id: return "This command only works in a Discord server." - - voice_channel = await adapter.get_user_voice_channel( - guild_id, event.source.user_id - ) + voice_channel = await adapter.get_user_voice_channel(guild_id, event.source.user_id) if not voice_channel: return "You need to be in a voice channel first." - - # Wire callbacks BEFORE join so voice input arriving immediately - # after connection is not lost. + # Wire callbacks BEFORE join so voice input arriving right after connection is not lost. self._bind_voice_input_callback(adapter) voice_profile = self._adapter_profile_for_source(event.source) if hasattr(adapter, "_on_voice_disconnect"): - adapter._on_voice_disconnect = functools.partial( - self._handle_voice_timeout_cleanup, adapter=adapter - ) - # Let the adapter's inactivity timer see the live voice-reply mode so it - # doesn't disconnect a deliberately text-only (/voice off) session. + cb = functools.partial(self._handle_voice_timeout_cleanup, adapter=adapter) + adapter._on_voice_disconnect = cb + # 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: logger.warning("Failed to join voice channel: %s", e) adapter._voice_input_callback = None err_lower = str(e).lower() - if "pynacl" in err_lower or "nacl" in err_lower or "davey" in err_lower: + if any(tok in err_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 success: - adapter._voice_text_channels[guild_id] = int(event.source.chat_id) - if hasattr(adapter, "_voice_sources"): - adapter._voice_sources[guild_id] = event.source.to_dict() - self._voice_mode[self._voice_key_for_source(event.source)] = "all" - self._save_voice_modes() - 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." - ) - # Join failed — clear callback - adapter._voice_input_callback = None - return "Failed to join voice channel. Check bot permissions (Connect + Speak)." + if not success: + adapter._voice_input_callback = None + return "Failed to join voice channel. Check bot permissions (Connect + Speak)." + adapter._voice_text_channels[guild_id] = int(event.source.chat_id) + if hasattr(adapter, "_voice_sources"): + adapter._voice_sources[guild_id] = event.source.to_dict() + self._voice_mode[self._voice_key_for_source(event.source)] = "all" + self._save_voice_modes() + 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." + ) 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 or not hasattr(adapter, "leave_voice_channel"): + 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) + ): return "Not in a voice channel." - - if not hasattr(adapter, "is_in_voice_channel") or not 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: @@ -270,11 +222,10 @@ class GatewayVoiceMixin: return "Left voice channel." def _handle_voice_timeout_cleanup(self, chat_id: str, *, adapter=None) -> None: - """Called by the adapter when a voice channel times out. + """Adapter callback on voice-channel timeout: clear runner-side voice_mode state. - Cleans up runner-side voice_mode state that the adapter cannot reach. ``adapter`` is the - Discord adapter that timed out (bound at join time); under multiplexing that is a - specific profile's bot, not necessarily ``self.adapters[DISCORD]``. + ``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]``. """ if adapter is None: adapter = self.adapters.get(Platform.DISCORD) @@ -284,40 +235,27 @@ class GatewayVoiceMixin: self._set_adapter_auto_tts_disabled(adapter, chat_id, disabled=True) def _is_duplicate_voice_transcript(self, guild_id: int, user_id: int, transcript: str) -> bool: - """Suppress repeated STT outputs for the same recent utterance. - - Voice capture can occasionally emit the same utterance twice a few seconds apart, which - creates a second queued agent run and overlapping spoken replies. + """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"\s+", " ", transcript).strip().lower() - normalized = re.sub(r"[^\w\s]", "", normalized) + normalized = re.sub(r"[^\w\s]", "", re.sub(r"\s+", " ", transcript).strip().lower()) if not normalized: return False - now = time.monotonic() - window_seconds = 12.0 key = (guild_id, user_id) recent_store = getattr(self, "_recent_voice_transcripts", None) if not isinstance(recent_store, dict): - recent_store = {} - self._recent_voice_transcripts = recent_store - recent = [ - (ts, txt) - for ts, txt in recent_store.get(key, []) - if now - ts <= window_seconds - ] - + 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: + if prior == normalized or ( + len(prior) >= 16 and len(normalized) >= 16 + and SequenceMatcher(None, prior, normalized).ratio() >= 0.95 + ): recent_store[key] = recent return True - if len(prior) >= 16 and len(normalized) >= 16: - if SequenceMatcher(None, prior, normalized).ratio() >= 0.95: - recent_store[key] = recent - return True - recent.append((now, normalized)) recent_store[key] = recent[-5:] return False @@ -333,137 +271,98 @@ class GatewayVoiceMixin: """ if adapter is None: adapter = self.adapters.get(Platform.DISCORD) - if not adapter: - return - - text_ch_id = adapter._voice_text_channels.get(guild_id) + text_ch_id = adapter._voice_text_channels.get(guild_id) if adapter else None if not text_ch_id: return - - # Build source — reuse the linked text channel's metadata when available - # so voice input shares the same session as the bound text conversation. + # 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 = str(user_id) - source.user_name = str(user_id) + 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", + 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), ) - - # Check authorization before processing voice input if not self._is_user_authorized(source): 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], + guild_id, user_id, transcript[:100], ) return - - # Show transcript in text channel (after auth, with mention sanitization) - try: + # Echo the transcript into the text channel (after auth, with mention sanitization). + with suppress(Exception): channel = adapter._client.get_channel(text_ch_id) if channel: safe_text = transcript[:2000].replace("@everyone", "@\u200beveryone").replace("@here", "@\u200bhere") await channel.send(f"**[Voice]** <@{user_id}>: {safe_text}") - except Exception: - pass - - # Build a synthetic MessageEvent for the normal pipeline; SimpleNamespace raw_message lets - # _get_guild_id() extract guild_id and _send_voice_reply() play audio in the voice channel. - from types import SimpleNamespace - # Resolve the bound text channel's channel_prompt so voice input gets - # the same per-channel context as typed messages (#50149). - channel_prompt: Optional[str] = None + # 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): - try: + with suppress(Exception): resolved = resolver(str(text_ch_id)) channel_prompt = resolved if isinstance(resolved, str) else None - except Exception: - channel_prompt = 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, + source=source, text=transcript, message_type=MessageType.VOICE, raw_message=SimpleNamespace(guild_id=guild_id, guild=None), 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, + 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 (dedup), or voice input + base adapter auto-TTS already - handled it (skip_double) — 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. + 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. """ if not response or response.startswith("Error:"): return False - chat_id = event.source.chat_id - voice_key = self._voice_key_for_source(event.source) - voice_mode = self._voice_mode.get(voice_key) - is_voice_input = (event.message_type == MessageType.VOICE) - + voice_mode = self._voice_mode.get(self._voice_key_for_source(event.source)) + is_voice_input = event.message_type == MessageType.VOICE adapter = self._adapter_for_source(event.source) adapter_auto_tts = False if adapter and hasattr(adapter, "_should_auto_tts_for_chat"): - try: + with suppress(Exception): adapter_auto_tts = bool(adapter._should_auto_tts_for_chat(chat_id)) - except Exception: - adapter_auto_tts = False - - should = ( - (voice_mode == "all") + # ``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) - # ``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. or (voice_mode is None and adapter_auto_tts) - ) - if not should: + ): 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, ) return False - - # Dedup: agent already called TTS tool in THIS turn only - last_user_idx = None - for i, msg in enumerate(reversed(agent_messages)): - if msg.get("role") == "user": - last_user_idx = len(agent_messages) - 1 - i; break - turn_messages = agent_messages[last_user_idx:] if last_user_idx is not None else agent_messages - has_agent_tts = any( - msg.get("role") == "assistant" - and any( - (tc.get("function") or {}).get("name") == "text_to_speech" - for tc in (msg.get("tool_calls") or []) - ) - for msg in turn_messages - ) - if has_agent_tts: + # 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 + if any( + (tc.get("function") or {}).get("name") == "text_to_speech" + for msg in turn_messages 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. @@ -483,12 +382,9 @@ class GatewayVoiceMixin: tts_text = _strip_markdown_for_tts(text) if not tts_text: return - - # Platforms whose native voice bubbles require Ogg/Opus (OPUS_VOICE_PLATFORMS — - # Telegram, Matrix, Feishu, WhatsApp, Signal) get an explicit .ogg path; the TTS tool's - # central container repair guarantees real Ogg/Opus bytes for every provider. + # 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 ) @@ -497,61 +393,44 @@ class GatewayVoiceMixin: except (json.JSONDecodeError, TypeError): logger.warning("Auto voice reply TTS returned invalid JSON: %s", result_json[:200] if result_json else result_json) return - - # Delivery may be one combined file or several separately valid files (combination - # unavailable or over a platform limit); preserve legacy single-file results. - actual_paths = result.get("file_paths") or [ - result.get("file_path", audio_path) - ] + # One combined file or several separately valid files (combination unavailable or + # over a platform limit); legacy single-file results keep working. actual_paths = [ - str(path) for path in actual_paths - if path and os.path.isfile(path) + 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 - - adapter = self._adapter_for_source(event.source) - - # If connected to a voice channel, play there instead of sending a file - guild_id = self._get_guild_id(event) - play_in_voice_channel = getattr(adapter, "play_in_voice_channel", None) - is_in_voice_channel = getattr(adapter, "is_in_voice_channel", None) - send_voice = getattr(adapter, "send_voice", None) - in_voice_channel = bool( - guild_id - and callable(play_in_voice_channel) - and callable(is_in_voice_channel) - and is_in_voice_channel(guild_id) - ) - reply_anchor = self._reply_anchor_for_event(event) - thread_meta = self._thread_metadata_for_source(event.source, reply_anchor) - if not in_voice_channel and callable(send_voice): - # Mark the auto voice reply as 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, not a silent message. Clone first so - # we don't mutate metadata shared with concurrent typing-indicator state. - if thread_meta is not None: - thread_meta = dict(thread_meta) - thread_meta["notify"] = True - else: - thread_meta = {"notify": True} - for actual_path in actual_paths: - if in_voice_channel: - play_voice = cast(Callable[..., Awaitable[Any]], play_in_voice_channel) - await play_voice(guild_id, actual_path) - elif callable(send_voice): - send_voice_call = cast(Callable[..., Awaitable[Any]], send_voice) - send_kwargs: Dict[str, Any] = { - "chat_id": event.source.chat_id, - "audio_path": actual_path, - "reply_to": reply_anchor, - "metadata": thread_meta, - } - await send_voice_call(**send_kwargs) + await self._deliver_voice_reply(event, actual_paths) except Exception as e: logger.warning("Auto voice reply failed: %s", e, exc_info=True) finally: for p in ({audio_path, *actual_paths} - {None}): with suppress(OSError): os.unlink(p) + + async def _deliver_voice_reply(self, event: MessageEvent, audio_paths: List[str]) -> None: + """Play the files in the connected voice channel, else send them as voice messages.""" + adapter = self._adapter_for_source(event.source) + 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): + 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. + 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, + ) diff --git a/gateway/run_watchers.py b/gateway/run_watchers.py index d0fc1e9821..678813d9a1 100644 --- a/gateway/run_watchers.py +++ b/gateway/run_watchers.py @@ -7,11 +7,11 @@ so ``patch("gateway.run.X")`` keeps intercepting them at call time. from __future__ import annotations -import logging -from typing import TYPE_CHECKING import asyncio +import logging import time -from typing import Any, Dict, Optional +from collections import Counter +from typing import TYPE_CHECKING, Any, Dict, Optional if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) from gateway.run import GatewayRunner, TurnRunner # noqa: F401 @@ -19,175 +19,155 @@ if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) # Log-record parity with the origin module. logger = logging.getLogger("gateway.run") +_MAX_FINALIZE_RETRIES = 3 +_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 + + +async def _interruptible_sleep(runner, seconds: int) -> None: + """Sleep in 1s increments so the watcher stops quickly when ``runner._running`` flips.""" + for _ in range(seconds): + if not runner._running: + break + await asyncio.sleep(1) + 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: runs ``on_session_finalize`` hooks, - cleans up the cached agent's tool resources, evicts the cache entry, and marks the session - finalized so it is not finalized again. + """Background task that finalizes expired sessions (``on_session_finalize`` hooks, cached + agent teardown, cache eviction, ``expiry_finalized`` flag) and runs the cache/store sweeps. """ - from gateway.run import _AGENT_PENDING_SENTINEL await asyncio.sleep(60) # initial delay — let the gateway fully start - _finalize_failures: dict[str, int] = {} # session_id -> consecutive failure count - _MAX_FINALIZE_RETRIES = 3 + finalize_failures: dict[str, int] = {} # session_id -> consecutive failure count while self._running: try: - await self.async_session_store._ensure_loaded() - # Collect expired sessions first, then log a single summary. - _expired_entries = [] - for key, entry in list(self.session_store._entries.items()): - if entry.expiry_finalized: - continue - if not await self.async_session_store._is_session_expired(entry): - continue - _expired_entries.append((key, entry)) - - if _expired_entries: - # Extract platform names from session keys for a compact summary. - # Keys look like "agent:main:telegram:dm:12345" — platform is field [2]. - _platforms: dict[str, int] = {} - for _k, _e in _expired_entries: - _parts = _k.split(":") - _plat = _parts[2] if len(_parts) > 2 else "unknown" - _platforms[_plat] = _platforms.get(_plat, 0) + 1 - _plat_summary = ", ".join( - f"{p}:{c}" for p, c in sorted(_platforms.items()) - ) + expired = await self._collect_expired_sessions() + if expired: + platforms = Counter(_platform_of_key(k, "unknown") for k, _ in expired) logger.info( "Session expiry: %d sessions to finalize (%s)", - len(_expired_entries), _plat_summary, + len(expired), ", ".join(f"{p}:{c}" for p, c in sorted(platforms.items())), ) - - for key, entry in _expired_entries: - try: - try: - _parts = key.split(":") - _platform = _parts[2] if len(_parts) > 2 else "" - # Off-loop + bounded: plugin finalize hooks can block arbitrarily, and - # this watcher runs on the gateway event loop. - await self._finalize_session_off_loop( - session_id=entry.session_id, - platform=_platform, - reason="session_expired", - ) - except Exception: - pass - # Close the cached agent's memory provider and tool resources. Idle agents - # live in _agent_cache (not _running_agents), so look there. - _cached_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) - _cached_agent = _cached[0] if isinstance(_cached, tuple) else _cached if _cached else None - # Fall back to _running_agents in case the agent is - # still mid-turn when the expiry fires. - if _cached_agent is None: - _exp_state = self._peek_session_state(key) - _cached_agent = _exp_state.turn.agent if _exp_state else None - if _cached_agent and _cached_agent is not _AGENT_PENDING_SENTINEL: - await self._cleanup_agent_resources_off_loop( - _cached_agent, context="session expiry" - ) - # Drop the cache entry so the AIAgent (LLM clients, tool schemas, memory - # provider refs) can be GC'd; otherwise the cache grows unbounded. - self._evict_cached_agent(key) - # Permanent finalization: one funnel call drops every conversation-scoped - # dict AND boundary security state so they don't grow unbounded. Idle - # agent-cache eviction must NOT do this — that session is still alive and a - # resumed turn rebuilds from these overrides. Only finalize, /new, /reset clear. - self._clear_conversation_scope( - key, reason="expiry_finalized" - ) - # Persist finalized flag (sessions.json AND state.db, single write-path); - # also drops the /model override — finalization is a conversation boundary. - await self.async_session_store.set_expiry_finalized(entry) - logger.debug( - "Session expiry finalized for %s", - entry.session_id, - ) - _finalize_failures.pop(entry.session_id, None) - except Exception as e: - failures = _finalize_failures.get(entry.session_id, 0) + 1 - _finalize_failures[entry.session_id] = failures - if failures >= _MAX_FINALIZE_RETRIES: - logger.warning( - "Session finalize gave up after %d attempts for %s: %s. " - "Marking as finalized to prevent infinite retry loop.", - failures, entry.session_id, e, - ) - await self.async_session_store.set_expiry_finalized( - entry, clear_model_override=False - ) - _finalize_failures.pop(entry.session_id, None) - else: - logger.debug( - "Session finalize failed (%d/%d) for %s: %s", - failures, _MAX_FINALIZE_RETRIES, entry.session_id, e, - ) - - if _expired_entries: - _done = sum( - 1 for _, e in _expired_entries if e.expiry_finalized - ) - _failed = len(_expired_entries) - _done - if _failed: + 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, + "Session expiry done: %d finalized, %d pending retry", done, failed ) else: - logger.info( - "Session expiry done: %d finalized", _done, - ) - - # Sweep agents idle beyond the TTL regardless of session reset policy: sessions with - # long / "never" reset windows would otherwise pin memory for the gateway's life. - try: - _idle_evicted = self._sweep_idle_cached_agents() - if _idle_evicted: - logger.info( - "Agent cache idle sweep: evicted %d agent(s)", - _idle_evicted, - ) - except Exception as _e: - logger.debug("Idle agent sweep failed: %s", _e) - - # Neither LRU cap nor idle TTL knows what a cached transcript costs in memory, so a - # busy gateway keeps every warm session's tool output resident until the RSS limit. - try: - 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_ts = getattr(self, "_last_session_store_prune_ts", 0.0) - _prune_interval = 3600.0 # once per hour - if time.time() - _last_prune_ts > _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, - ) - except Exception as _e: - logger.debug("SessionStore prune failed: %s", _e) - self._last_session_store_prune_ts = time.time() + logger.info("Session expiry done: %d finalized", done) + await self._expiry_housekeeping() except Exception as e: logger.debug("Session expiry watcher error: %s", e) - # Sleep in small increments so we can stop quickly - for _ in range(interval): - if not self._running: - break - await asyncio.sleep(1) + await _interruptible_sleep(self, interval) + + 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 + + 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. + """ + 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: + 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, + ) + store = self.async_session_store + await store.set_expiry_finalized(entry, clear_model_override=False) + failures.pop(sid, None) + + 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. + await self._finalize_session_off_loop( + session_id=entry.session_id, + platform=_platform_of_key(key), + reason="session_expired", + ) + 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 + 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. + self._evict_cached_agent(key) + self._clear_conversation_scope(key, reason="expiry_finalized") + await self.async_session_store.set_expiry_finalized(entry) + logger.debug("Session expiry finalized for %s", entry.session_id) + + 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. + try: + if evicted := self._sweep_idle_cached_agents(): + logger.info("Agent cache idle sweep: evicted %d agent(s)", evicted) + except Exception as e: + logger.debug("Idle agent sweep failed: %s", e) + # Neither LRU cap nor idle TTL knows what a cached transcript costs in memory. + try: + 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: + 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) + except Exception as e: + logger.debug("SessionStore prune failed: %s", e) + self._last_session_store_prune_ts = time.time() def _session_stall_timeout_seconds(self) -> float: """Return configured stall timeout (seconds); 0 disables the watchdog.""" @@ -195,30 +175,18 @@ class GatewaySessionWatchersMixin: return _float_env("HERMES_SESSION_STALL_TIMEOUT", 300) def _iter_gateway_adapters(self): - """Yield every live platform adapter (default + multiplex profiles).""" + """Yield every live platform adapter (default + multiplex profiles), deduped by identity.""" seen: set[int] = set() - for adapter in list(getattr(self, "adapters", {}).values()): - if adapter is None: - continue - aid = id(adapter) - if aid in seen: - continue - seen.add(aid) - yield adapter - for amap in list(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 None: - continue - aid = id(adapter) - if aid in seen: - continue - seen.add(aid) - yield adapter + 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]: """Return the shared activity snapshot for stall progress: the single source is - ``AIAgent.get_activity_summary()`` / ``agent.session_activity``; no turn-start or - pending-inbound clocks. + ``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) @@ -232,13 +200,35 @@ class GatewaySessionWatchersMixin: return None 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).""" + 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()): + 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: + 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.run import _STALL_NOTIFY_SEND_TIMEOUT_SECONDS 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, @@ -246,171 +236,105 @@ class GatewaySessionWatchersMixin: notified_map = getattr(self, "_session_stall_notified", None) if notified_map is None: - notified_map = {} - self._session_stall_notified = notified_map - + notified_map = self._session_stall_notified = {} sent = 0 now = time.time() - 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()): - 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 - event = overflow[0] - source = getattr(event, "source", None) - adapter = ( - self._adapter_for_source(source) if source is not None else None - ) - if adapter is None: - continue - candidates[session_key] = (adapter, event) - + candidates = 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()): - has_pending = pending_event is not None - activity = ( - self._session_activity_for_stall(session_key) if has_pending else None - ) - idle_seconds = ( - resolve_session_idle_seconds_from_activity(activity, now=now) - if has_pending - else None - ) + 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=has_pending, + timeout_seconds=timeout_seconds, idle_seconds=idle_seconds, has_pending_inbound=True ): notified_map.pop(session_key, None) already = False - if not should_emit_session_stall_notification( - timeout_seconds=timeout_seconds, - idle_seconds=idle_seconds, - has_pending_inbound=has_pending, - already_notified=already, + 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, ): continue - - if idle_seconds is None: - continue - mins = max(1, int(idle_seconds // 60)) - activity = activity or {} - 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, - mins, - 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) - chat_id = getattr(source, "chat_id", None) if source is not None else None - if not chat_id: - logger.warning( - "Session stall notify skipped (no chat_id): session=%s", - session_key, - ) - # Cannot deliver; latch to avoid log spam every tick. - notified_map[session_key] = True - continue - # Re-read pending state + activity IMMEDIATELY before delivery: the snapshot above ages - # while earlier candidates await sends; an agent that progressed (or drained its queue) - # must not get a false stall notice. Abort, latch un-set, so the next tick re-evaluates. - 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 + if await self._notify_session_stall( + session_key, adapter, pending_event, idle_seconds, activity or {}, + timeout_seconds, notified_map, ): - logger.info( - "Session stall notify aborted (no longer stale): " - "session=%s pending=%s fresh_idle=%s", - session_key, - still_pending, - fresh_idle, - ) - # Re-arm: drop any stale latch so a FUTURE genuine stall - # episode notifies again. - notified_map.pop(session_key, None) - continue - try: - metadata = ( - self._thread_metadata_for_source(source) - if source is not None and hasattr(self, "_thread_metadata_for_source") - else None - ) - # 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. - try: - result = await asyncio.wait_for( - adapter.send( - str(chat_id), - format_session_stall_notification(idle_seconds), - 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, - ) - continue # do not latch; retry next tick - # 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"), - ) - continue # do not latch; retry next tick sent += 1 - notified_map[session_key] = True - except Exception as exc: - logger.warning( - "Session stall notify failed for %s: %s", - session_key, - exc, - ) - # Do not latch — retry next watcher tick until delivery or episode clear. - # 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) - 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.""" + from gateway.run import _STALL_NOTIFY_SEND_TIMEOUT_SECONDS + from gateway.session_stall import ( + format_session_stall_notification, + 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 (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) + # 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. + try: + result = await asyncio.wait_for( + adapter.send(str(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 + # 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 + 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 + 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. @@ -432,11 +356,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 ``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``). + """Periodic 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``). """ # Short initial delay so startup reconnect noise does not false-fire. await asyncio.sleep(min(30.0, max(1.0, float(interval)))) @@ -447,9 +369,4 @@ class GatewaySessionWatchersMixin: await self._check_session_stalls(timeout) except Exception as exc: logger.debug("Session stall watcher error: %s", exc) - # Interruptible sleep - steps = max(1, int(float(interval))) - for _ in range(steps): - if not self._running: - break - await asyncio.sleep(1) + await _interruptible_sleep(self, max(1, int(float(interval)))) diff --git a/gateway/sticker_cache.py b/gateway/sticker_cache.py index f336e8fab9..40b79d7a18 100644 --- a/gateway/sticker_cache.py +++ b/gateway/sticker_cache.py @@ -4,6 +4,7 @@ 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 @@ -11,7 +12,6 @@ import time from typing import Optional from hermes_cli.config import get_hermes_home -import contextlib CACHE_PATH = get_hermes_home() / "sticker_cache.json" @@ -24,12 +24,10 @@ STICKER_VISION_PROMPT = ( def _load_cache() -> dict: - if CACHE_PATH.exists(): - try: - return json.loads(CACHE_PATH.read_text(encoding="utf-8")) - except (json.JSONDecodeError, OSError): - return {} - return {} + try: + return json.loads(CACHE_PATH.read_text(encoding="utf-8")) + except (FileNotFoundError, json.JSONDecodeError, OSError): + return {} def _save_cache(cache: dict) -> None: @@ -65,11 +63,9 @@ def build_sticker_injection(description: str, emoji: str = "", set_name: str = " ``[The user sent a sticker 😀 from "MyPack"~ It shows: "A cat waving" (=^.w.^=)]``. ``set_name`` is only shown together with an emoji. """ - context = "" + context = f" {emoji}" if emoji else "" if set_name and emoji: - context = f" {emoji} from \"{set_name}\"" - elif emoji: - context = f" {emoji}" + context += f" from \"{set_name}\"" return f"[The user sent a sticker{context}~ It shows: \"{description}\" (=^.w.^=)]" diff --git a/gateway/stream_consumer.py b/gateway/stream_consumer.py index a18e5021b7..00a08b9b07 100644 --- a/gateway/stream_consumer.py +++ b/gateway/stream_consumer.py @@ -1,9 +1,9 @@ """Gateway streaming consumer — bridges sync agent callbacks to async platform delivery. The agent fires stream_delta_callback(text) synchronously from its worker thread; -on_delta() queues it (queue.Queue) and the async run() task buffers, rate-limits, -and progressively edits a single platform message (send, then editMessageText — -supported everywhere; draft/native transports are optional per adapter). +on_delta() queues it and the async run() task buffers, rate-limits and progressively +edits a single platform message (send, then editMessageText — supported everywhere; +draft/native transports are optional per adapter). Credit: jobless0x (#774, #1312), OutThisLife (#798), clicksingh (#697). """ @@ -12,12 +12,14 @@ from __future__ import annotations import asyncio import concurrent.futures +import contextlib import inspect import logging import queue import secrets import threading import time +import uuid from dataclasses import dataclass from typing import Any, Callable, Optional @@ -39,35 +41,31 @@ from gateway.stream_consumer_fences import ( # noqa: F401 (re-exported) from gateway.stream_consumer_transport import StreamTransportMixin from gateway.stream_consumer_fallback import StreamFallbackMixin from gateway.stream_consumer_think import StreamThinkFilterMixin -import contextlib logger = logging.getLogger("gateway.stream_consumer") -# Queue sentinels (see run() for handling). -_DONE = object() # stream complete -_NEW_SEGMENT = object() # finalize current message, start a fresh one -_COMMENTARY = object() # (_COMMENTARY, text): completed interim commentary -_TOOL_PROGRESS = object() # (_TOOL_PROGRESS, line): overlay line for the native bubble -# (_FINAL_TEXT, text): authoritative completed final_response (incl. post-stream -# augmentation such as verifier footers) enqueued just before _DONE, so the -# finalize/seal delivers the TRUE final and the recorded payload reconciles. +# Queue sentinels (see _drain_queue()). Bare: _DONE (stream complete), _NEW_SEGMENT +# (finalize current message, start a fresh one), _REOPEN_SEED (EAGER native re-seed +# right after a clarify answer — WeCom typing is driven by the seed frame; lazy +# re-seed measured 48s of dead air). Tuple-shaped: (_COMMENTARY, text) completed +# interim commentary; (_TOOL_PROGRESS, line) overlay for the native bubble; +# (_FINAL_TEXT, text) authoritative completed final_response incl. post-stream +# augmentation, enqueued just before _DONE so the finalize delivers the TRUE final; +# (_FLUSH, threading.Event) barrier set once everything buffered is delivered; +# (_APPROVAL_BOUNDARY, future, cancelled_flag) approval/clarify prompt boundary. +_DONE = object() +_NEW_SEGMENT = object() +_COMMENTARY = object() +_TOOL_PROGRESS = object() _FINAL_TEXT = object() -# (_FLUSH, threading.Event): barrier — drain loop finalizes/delivers anything -# buffered, then sets the event. Used by flush_pending_sync() before a blocking -# interactive prompt so the prompt lands below buffered prose, not above it. _FLUSH = object() -# Interaction boundary (approval OR clarify prompt): finalize the current -# stream; post-prompt output goes via send() or a re-opened stream. _APPROVAL_BOUNDARY = object() -# EAGER native re-seed right after a clarify answer, before the first post-answer -# delta. On WeCom typing is driven by the seed frame (send_typing is a no-op) -# and lazy re-seed on first delta measured 48s of dead air. _REOPEN_SEED = object() -# Boundary finalize text when nothing has accumulated yet; overridable per -# boundary via close_for_approval_prompt(placeholder=...). +# Boundary finalize text when nothing has accumulated yet (overridable per boundary). _DEFAULT_BOUNDARY_PLACEHOLDER = "⏸ 等待审批中..." + @dataclass class StreamConsumerConfig: """Runtime config for a single stream consumer instance.""" @@ -75,21 +73,18 @@ class StreamConsumerConfig: buffer_threshold: int = _DEFAULT_STREAMING_BUFFER_THRESHOLD cursor: str = _DEFAULT_STREAMING_CURSOR buffer_only: bool = False - # When >0, deliver the final as a fresh message if the preview has been - # visible at least this long, so the visible timestamp reflects completion - # rather than first token. 0 = always edit in place. Enabled per-platform. + # >0: final goes out as a fresh message once the preview has been visible this + # long (timestamp reflects completion); 0 = always edit in place. fresh_final_after_seconds: float = 0.0 - # "auto"/"draft": native draft streaming when adapter+chat support it, else - # edit. "edit": progressive editMessageText. "off": handled by the gateway - # before the consumer is built. + # "auto"/"draft": native drafts when adapter+chat support it, else "edit" + # (progressive editMessageText). "off" is handled by the gateway. transport: str = "edit" - # Originating chat type ("dm", "group", ...); gates draft streaming, which - # is platform-specific (Telegram drafts are DM-only). - chat_type: str = "" + chat_type: str = "" # originating chat type; gates platform-specific drafts + @dataclass class _Tick: - """Everything one drain of the queue decided (flag locals of the old run() loop).""" + """Everything one drain of the queue decided.""" got_done: bool = False got_segment_break: bool = False got_flush: bool = False @@ -106,6 +101,7 @@ class _Tick: """Mid-stream tick: not finalizing, not a segment break, no commentary.""" return not self.got_done and not self.got_segment_break and self.commentary_text is None + class GatewayStreamConsumer( StreamTransportMixin, StreamFallbackMixin, @@ -113,28 +109,19 @@ class GatewayStreamConsumer( ): """Async consumer that progressively edits a platform message with streamed tokens. - Usage:: - - consumer = GatewayStreamConsumer(adapter, chat_id, config, metadata=metadata) - agent = AIAgent(..., stream_delta_callback=consumer.on_delta) - task = asyncio.create_task(consumer.run()) - # ... run agent in thread pool ... - consumer.finish() # signal completion - await task # wait for final edit + Usage: ``agent.stream_delta_callback = consumer.on_delta``; ``task = + create_task(consumer.run())``; after the agent finishes ``consumer.finish()`` + then ``await task`` for the final edit. """ - # Consecutive flood-control failures before progressive edits are disabled - # for the rest of the stream. + # Consecutive flood-control failures before progressive edits are disabled. _MAX_FLOOD_STRIKES = 3 - # Class-wide monotonic counter for draft ids (Telegram animates a draft only - # when the same non-zero draft_id is reused within a response). Seeded from - # a RANDOM nonce, not 0 or the clock: draft_id is the wire identity for the - # relay connector's per-(channel, draft_id) sealed-stream tombstones, which - # outlive this (scale-to-zero) process — a replayed id is answered out of - # the OLD tombstone and the reply is silently dropped. Epoch-ms seeds still - # collide on same-ms starts/forks/clock steps; 49 bits keeps ids + turn - # counts inside the connector's JS number range (2^53). + # Class-wide monotonic draft-id counter (Telegram animates a draft only when the + # same non-zero draft_id is reused). RANDOM seed: draft_id keys the relay + # connector's sealed-stream tombstones, which outlive this (scale-to-zero) + # process — a replayed id is answered out of the OLD tombstone and silently + # dropped. 49 bits keeps ids + turn counts inside JS number range (2^53). _draft_id_counter: int = secrets.randbits(49) def __init__( @@ -152,124 +139,103 @@ class GatewayStreamConsumer( self.chat_id = chat_id self.cfg = config or StreamConsumerConfig() self.metadata = metadata - # Fired whenever a fresh content bubble is created (first send, - # commentary, overflow chunk, fallback continuation) so the gateway - # opens the next tool-progress bubble BELOW it. Exceptions swallowed. + # Gateway hooks (exceptions swallowed): on_new_message whenever a fresh content + # bubble is created (next tool-progress bubble goes BELOW it); on_before_finalize + # once on entering finalization (pause typing refreshes). self._on_new_message = on_new_message - # Fired once on entering finalization so the gateway can pause typing - # refreshes before a slow rich-text final edit. self._on_before_finalize = on_before_finalize self._initial_reply_to_id = initial_reply_to_id - # Per-turn id for adapter.send_stream_frame() so concurrent consumers - # (/background, parallel subagents) don't interfere. - import uuid - self._turn_id = str(uuid.uuid4()) + self._turn_id = str(uuid.uuid4()) # keys send_stream_frame() per concurrent consumer # Returns False after /new or /stop; run() then abandons the stream. self._run_still_current = run_still_current or (lambda: True) - # Only platforms needing an explicit finalize call (DingTalk AI Cards) - # force a redundant final edit; ``is True`` keeps MagicMock adapters out. + # Only platforms needing an explicit finalize call (DingTalk AI Cards) force a + # redundant final edit; ``is True`` keeps MagicMock adapters out. self._adapter_requires_finalize: bool = ( getattr(adapter, "REQUIRES_EDIT_FINALIZE", False) is True ) - # Telegram bounds edit retries at 5s; a final-delivery fallback must not - # hold the stream task through a longer flood cooldown. + # Telegram bounds edit retries at 5s; a fallback must not wait longer. self._max_fallback_flood_retry_seconds = 5.0 self._queue: queue.Queue = queue.Queue() - self._accumulated = "" - # Mirror of ``_accumulated`` NOT truncated when overflow splits seal - # head chunks; records a reconcilable turn-final payload for splits. - self._stream_ledger = "" - self._message_id: Optional[str] = None - # monotonic() when ``_message_id`` was first assigned (fresh-final age). - self._message_created_ts: Optional[float] = None - # Every real preview id on screen this response (first send + - # continuations): fresh-final deletes them all so a reply split at the - # edit limit leaves no stale fragments above the final. The segment - # set holds only the active text segment: failure recovery must never - # delete an earlier finalized preamble/commentary message. + # Every real preview id on screen this response (fresh-final deletes them all); + # the per-segment set (see _reset_message_state) holds only the active segment + # so failure recovery never deletes an earlier finalized preamble/commentary. self._preview_message_ids: "set[str]" = set() - self._segment_preview_message_ids: "set[str]" = set() self._already_sent = False self._edit_supported = True # False once progressive edits stop working self._last_edit_time = 0.0 - self._last_sent_text = "" # skip redundant edits - # Most recent _send_or_edit split across continuation messages (the - # adapter adopted a new message id). - self._last_edit_overflowed = False - self._fallback_final_send = False - self._fallback_prefix = "" - # Fallback sends only the missing tail after a partial overflow - # delivery: the visible prefix is content, not a stale preview. - self._fallback_preserve_partial_messages = False - self._flood_strikes = 0 # consecutive flood-control edit failures + self._last_edit_overflowed = False # last _send_or_edit split into continuations + self._flood_strikes = 0 self._current_edit_interval = self.cfg.edit_interval # adaptive backoff - self._final_response_sent = False - # Final content reached the user even if the cosmetic final edit - # (cursor removal) then failed. - self._final_content_delivered = False - # Exact cleaned payload of the turn-final delivery that set the flags - # above; the gateway compares it to the completed final_response before - # trusting the flags (a successful finalize edit may carry only a stale - # preview snapshot). ``None`` = no record → legacy trust. - self._delivered_final_text: Optional[str] = None - # Answer delivered across multiple sealed messages (overflow split / - # continuation adoption). Payload-less split delivery must NOT inherit - # legacy trust (it swallowed complete replies after partial splits). - self._turn_split_delivery = False - # A full-final send timed out in a way that MAY have reached the - # platform — the only payload-less case that keeps legacy trust, since - # re-sending risks a duplicate rather than recovering a loss. - self._delivery_ambiguous = False self._delivered_commentary_texts: list[str] = [] - # Finalized visible text per segment, so has_delivered_text still - # matches after _reset_segment_state clears _last_sent_text. - self._delivered_segment_texts: list[str] = [] - # Think-block filter state (mirrors CLI's _stream_delta tag suppression). - self._in_think_block = False + self._delivered_segment_texts: list[str] = [] # finalized text per past segment + self._in_think_block = False # think-tag filter state (mirrors CLI _stream_delta) self._think_buffer = "" self._before_finalize_notified = False + self._reset_message_state() - # Transports, resolved at the start of run(). Draft: animated frames - # via adapter.send_draft instead of edits; the final still uses the - # first-send path (drafts have no message_id); the first draft failure - # disables drafts for the response. Native (WeCom msgtype "stream"): - # the ONLY delivery channel — seed, cumulative updates and finish=true - # all go through send_stream_frame(); any failure falls back to edit/send. + # Transports, resolved at the start of run(). Draft: animated frames via + # adapter.send_draft; the final still uses the first-send path; the first + # draft failure disables drafts. Native (WeCom msgtype "stream"): the ONLY + # delivery channel — any failure falls back to edit/send. self._use_draft_streaming = False self._draft_id: Optional[int] = None self._draft_failures = 0 self._use_native_streaming = False - # Seed frame sent (zero visible content but the bubble is open); the - # fallback decides from this whether the stream must be finalized first. - self._native_stream_opened = False - # Visible chars last pushed; throttles under WeCom's 30 frames/min. - self._native_last_pushed_len = 0 - # Boundary state from close_for_approval_prompt(); race-free because - # boundaries are processed serially. ``_boundary_reopen`` keeps native - # enabled so post-prompt output re-opens a fresh stream via the lazy - # re-seed (clarify: short waits) instead of degrading to send() - # (approval: unbounded waits, stream may go stale). + self._native_stream_opened = False # seed sent: bubble open, zero content + self._native_last_pushed_len = 0 # throttle under WeCom's 30 frames/min + # Boundary state from close_for_approval_prompt() (race-free: boundaries are + # processed serially). reopen=True (clarify, short waits) keeps native enabled + # so post-prompt output re-opens a fresh stream; approval degrades to send(). self._boundary_placeholder = _DEFAULT_BOUNDARY_PLACEHOLDER self._boundary_reason = "Approval" self._boundary_reopen = False - # Reopen requested but nothing re-seeded yet: got_done must not open a - # stream just to emit a lone "✅". An EAGER re-seed (_REOPEN_SEED) - # already opened a fresh bubble before any content: got_done must - # actively finalize it or a blank typing bubble hangs forever. + # Reopen requested but nothing re-seeded: got_done must not open a stream just + # to emit a lone "✅". An EAGER re-seed already opened a bubble: got_done must + # actively finalize it or a blank typing bubble hangs. self._awaiting_reopen_after_boundary = False self._reopen_seeded_eagerly = False + + def _reset_message_state(self) -> None: + """Per-message (segment) state: fresh at construction and after each segment break.""" + self._message_id: Optional[str] = None + self._message_created_ts: Optional[float] = None # fresh-final age + # ``_stream_ledger`` mirrors ``_accumulated`` but is NOT truncated when + # overflow splits seal head chunks (reconcilable turn-final payload). + self._accumulated = self._stream_ledger = "" + self._last_sent_text = "" # skip redundant edits + self._fallback_final_send = False + self._fallback_prefix = "" + # Fallback sends only the missing tail after a partial overflow delivery. + self._fallback_preserve_partial_messages = False + self._segment_preview_message_ids: "set[str]" = set() # Tool-progress overlay (native only): shown in the bubble until text arrives. self._tool_progress_lines: list[str] = [] self._tool_progress_active: bool = False + self._clear_turn_final_flags() + + def _clear_turn_final_flags(self) -> None: + """Reset every turn-final delivery flag to "nothing delivered yet". + + ``_delivered_final_text``: exact cleaned payload of the turn-final delivery; + the gateway compares it to the completed final_response before trusting the + flags (a successful finalize edit may carry only a stale preview); None = + legacy trust. ``_turn_split_delivery``: answer delivered across multiple + sealed messages — payload-less split delivery must NOT inherit legacy trust. + ``_delivery_ambiguous``: a full-final send timed out in a way that MAY have + reached the platform — the only payload-less case that keeps legacy trust. + """ + self._final_response_sent = False + self._final_content_delivered = False # content landed even if the cosmetic edit failed + self._delivered_final_text: Optional[str] = None + self._turn_split_delivery = False + self._delivery_ambiguous = False def _stream_is_message(self) -> bool: """Whether THIS chat's transport treats the stream as the message. - Prefers the adapter's per-chat probe (a multi-platform relay adapter's - class attribute can only reflect its primary identity), falling back to - the legacy attribute. Both are resolved on the CLASS to stay - MagicMock-safe (auto-created instance attributes are truthy). + Per-chat probe first (a relay adapter's class attribute only reflects its + primary identity), else the legacy attribute; both on the CLASS (MagicMock-safe). """ probe = getattr(type(self.adapter), "stream_is_message_for_chat", None) if callable(probe): @@ -291,26 +257,15 @@ class GatewayStreamConsumer( def _compose_frame_content(self) -> str: """Native frame content: text, with any tool-progress lines below a rule.""" - if self._accumulated and self._tool_progress_lines: - return self._accumulated + "\n\n---\n" + "\n".join(self._tool_progress_lines) - elif self._accumulated: - return self._accumulated - elif self._tool_progress_lines: - return "\n".join(self._tool_progress_lines) - return "" + progress = "\n".join(self._tool_progress_lines) + return "\n\n---\n".join(p for p in (self._accumulated, progress) if p) - def _metadata_for_send( - self, - *, - final: bool = False, - expect_edits: bool = False, - ) -> dict | None: - """Per-send metadata for stream-created messages. + def _metadata_for_send(self, *, final: bool = False, expect_edits: bool = False) -> dict | None: + """Per-send metadata. - ``final`` sets notify=True (Mattermost treats notify-worthy sends as - final content when deciding whether a broken thread root may fall back - flat). ``expect_edits`` keeps editable previews on Telegram's legacy - send path while final sends may use richer delivery. + ``final`` → notify=True (Mattermost treats notify-worthy sends as final when a + broken thread root may fall back flat); ``expect_edits`` keeps editable + previews on Telegram's legacy send path. """ meta = dict(self.metadata) if self.metadata else {} if self._initial_reply_to_id: @@ -321,25 +276,13 @@ class GatewayStreamConsumer( meta["notify"] = True return meta or None - @property - def already_sent(self) -> bool: - """True if at least one message was sent or edited during the run.""" - return self._already_sent - - @property - def final_response_sent(self) -> bool: - """True when the stream consumer delivered the final assistant reply.""" - return self._final_response_sent - - @property - def message_id(self) -> str | None: - """Message ID of the last-sent or edited message.""" - return self._message_id - - @property - def final_content_delivered(self) -> bool: - """True when the final content reached the user, even if the cosmetic final edit failed.""" - return self._final_content_delivered + # Read-only views for the gateway: a message was sent/edited; the final reply was + # delivered; id of the last-sent/edited message; final content reached the user + # even if the cosmetic final edit failed. + already_sent = property(lambda self: self._already_sent) + final_response_sent = property(lambda self: self._final_response_sent) + message_id = property(lambda self: self._message_id) + final_content_delivered = property(lambda self: self._final_content_delivered) async def _notify_before_finalize(self) -> None: """Run the pre-finalize hook exactly once, swallowing hook errors.""" @@ -348,19 +291,16 @@ class GatewayStreamConsumer( self._before_finalize_notified = True if self._on_before_finalize is None: return - try: + with contextlib.suppress(Exception): result = self._on_before_finalize() if inspect.isawaitable(result): await result - except Exception: - pass def _append_accumulated(self, text: str) -> None: """Append to the live buffer and the split-stable stream ledger.""" if not text: return - # Real text overwrites the tool-progress overlay. - if self._tool_progress_lines: + if self._tool_progress_lines: # real text overwrites the overlay self._tool_progress_lines.clear() self._tool_progress_active = False self._accumulated += text @@ -369,10 +309,9 @@ class GatewayStreamConsumer( def _mark_skip_redundant_finalize(self) -> None: """Mark the turn final as delivered by a prior mid-stream edit. - Records what was ACKED on the wire, not ``_accumulated``: a throttled - edit stream can reach this state with the last acked edit still holding - an earlier (cursor-suffixed) preview, and recording the accumulator - would let that frozen preview suppress the corrective send. + Records what was ACKED on the wire, not ``_accumulated``: a throttled stream's + last ack may be an older cursor-suffixed preview, which must not suppress the + corrective send. """ self._mark_final_delivered() acked = self._last_sent_text or self._accumulated @@ -387,55 +326,44 @@ class GatewayStreamConsumer( if record is not None: self._record_turn_final_payload(record) + def _display_payload(self, text: str) -> str: + """Normalize like ``_send_or_edit`` output: directive strip + fence close + strip.""" + return ensure_closed_code_fences(self._clean_for_display(text or "")).strip() + def _record_turn_final_payload(self, text: str) -> None: """Record what the user actually saw as this turn's final answer. - Normalized like ``_send_or_edit`` output (media-directive strip + fence - closing) so the gateway can compare it to the completed final_response. - On a multi-message split ``text`` is only the trailing chunk (overflow - truncates ``_accumulated`` as head chunks seal), so the un-truncated - ``_stream_ledger`` is recorded instead — else the gateway sees a - mismatch and re-sends an answer the user already received. + On a split ``text`` is only the trailing chunk, so the un-truncated + ``_stream_ledger`` is recorded instead — else the gateway sees a mismatch and + re-sends an answer the user already received. """ source = text or "" if self._turn_split_delivery and self._stream_ledger: source = self._stream_ledger - self._delivered_final_text = ensure_closed_code_fences( - self._clean_for_display(source) - ).strip() + self._delivered_final_text = self._display_payload(source) def delivered_final_matches(self, final_text: str) -> Optional[bool]: """Tri-state reconcile of the recorded turn-final payload against ``final_text``. - A *successful* finalize edit can still carry only a stale preview, so - call success alone must not confirm delivery. - True: recorded payload (or an earlier segment/commentary) matches — - suppressing the normal final send is safe. False: recorded payload - differs, or payload-less multi-message split (flag alone must not - suppress). None: nothing recorded on a non-split legacy/ambiguous path; - caller keeps its flag-trusting behavior. + A *successful* finalize edit can still carry only a stale preview, so call + success alone must not confirm delivery. True: recorded payload (or an + earlier segment/commentary) matches. False: payload differs, or payload-less + split. None: nothing recorded on a legacy/ambiguous path (caller trusts flags). """ - target = ensure_closed_code_fences( - self._clean_for_display(final_text or "") - ).strip() + target = self._display_payload(final_text) if not target: return None if self._delivered_final_text is None: if self._turn_split_delivery: return False - # No recorded payload: judge against the FINAL content rather than - # trusting the flag — a consumer whose visible text lacks the - # completed response (first-edit prefix, mid-stream truncation) has + # No recorded payload: judge against the FINAL content, not the flag — a + # consumer whose visible text lacks the completed response has # demonstrably NOT delivered it. ``_already_sent`` gates the match: - # draft frames set ``_last_sent_text`` but are ephemeral and - # deliberately don't set ``_already_sent``. + # draft frames set ``_last_sent_text`` but deliberately not ``_already_sent``. if self._already_sent and self.has_delivered_text(final_text): return True - # Only a timed-out full-final send that MAY have landed keeps legacy - # trust — re-sending risks a duplicate. - if self._delivery_ambiguous: - return None - return False + # Only a timed-out full-final send that MAY have landed keeps legacy trust. + return None if self._delivery_ambiguous else False if self._delivered_final_text.strip() == target: return True # A segment break / commentary may have delivered it under another record. @@ -446,8 +374,7 @@ class GatewayStreamConsumer( target = self._clean_for_display(text or "").strip() if not target: return False - visible_prefix = self._visible_prefix().strip() - if visible_prefix == target: + if self._visible_prefix().strip() == target: return True return any( sent.strip() == target @@ -466,18 +393,12 @@ class GatewayStreamConsumer( ) -> asyncio.Future: """Queue an interaction boundary (approval / clarify prompt) from sync context. - run() processes it serially: finalize the current native stream with - accumulated text (``placeholder`` when nothing accumulated yet), then - per ``reopen``: False (approval; long unbounded waits, stream may go - stale) disables native streaming and batches post-prompt output into - one send(); True (clarify) keeps native enabled so post-prompt output - re-opens a fresh stream via the lazy re-seed, degrading to send() if - that fails. ``reason`` labels log lines ("Approval"/"Clarify"). - - Returns a (Future, cancelled_flag) tuple; the Future resolves True once - processed. cancelled_flag is kept for callers that set it on timeout - but the handler no longer reads it (finalize always runs). Without - native streaming returns a bare, already-resolved Future. + run() finalizes the current native stream (``placeholder`` when empty), then + per ``reopen``: False (approval; unbounded waits) degrades to one send() at + got_done; True (clarify) keeps native enabled so post-prompt output re-opens + a fresh stream. Returns (Future, cancelled_flag); the Future resolves True + once processed (cancelled_flag is legacy, no longer read). Without native + streaming returns a bare, already-resolved Future. """ loop = None with contextlib.suppress(RuntimeError): @@ -494,7 +415,6 @@ class GatewayStreamConsumer( self._boundary_reopen = bool(reopen) boundary_future = loop.create_future() if loop else concurrent.futures.Future() - cancelled_flag = {"cancelled": False} self._queue.put((_APPROVAL_BOUNDARY, boundary_future, cancelled_flag)) return boundary_future, cancelled_flag @@ -507,11 +427,8 @@ class GatewayStreamConsumer( def flush_pending_sync(self, timeout: float = 5.0) -> bool: """Block the agent worker thread until everything queued so far is delivered. - Enqueues a ``(_FLUSH, Event)`` barrier; run() drains earlier items - (FIFO), finalizes the current segment, then sets the event. Returns - False on timeout so the caller proceeds even if the consumer task is - not running. Used before a blocking interactive prompt (clarify poll) - so the question doesn't land ABOVE its own explanation. + ``(_FLUSH, Event)`` barrier: run() drains earlier items (FIFO), finalizes the + segment, sets the event. False on timeout (consumer task may not be running). """ evt = threading.Event() try: @@ -520,28 +437,29 @@ class GatewayStreamConsumer( return False return evt.wait(timeout=max(0.0, float(timeout))) - def request_reopen_seed(self) -> None: - """Thread-safe: request an EAGER native re-seed after a clarify answer. - - Posts _REOPEN_SEED so run() sends an empty seed frame before the first - post-answer delta (WeCom typing bubble reappears immediately). No-op - unless reopen-pending on a native stream with no stream open, so a - stray call can't open a spurious bubble mid-stream or on approval. - """ - if ( + def _reopen_seed_pending(self) -> bool: + """Native stream, reopen requested after a boundary, nothing open yet.""" + return ( self._use_native_streaming and self._awaiting_reopen_after_boundary and not self._native_stream_opened - ): + ) + + def request_reopen_seed(self) -> None: + """Thread-safe: request an EAGER native re-seed after a clarify answer. + + No-op unless reopen-pending on a native stream with no stream open, so a + stray call can't open a spurious bubble mid-stream or on approval. + """ + if self._reopen_seed_pending(): self._queue.put(_REOPEN_SEED) def _notify_new_message(self) -> None: """Fire the on_new_message callback, swallowing any errors.""" - cb = self._on_new_message - if cb is None: + if self._on_new_message is None: return try: - cb() + self._on_new_message() except Exception: logger.debug("on_new_message callback error", exc_info=True) @@ -549,14 +467,12 @@ class GatewayStreamConsumer( def _signal_flush(flush_event) -> None: """Wake a thread blocked in flush_pending_sync(), swallowing errors. - Every loop path that consumed a ``_FLUSH`` barrier (including early - ``continue`` paths) must call this; a missed set isn't a deadlock - (bounded wait) but stalls the caller for the full timeout. + Every loop path that consumed a ``_FLUSH`` barrier (incl. early ``continue``) + must call this; a missed set stalls the caller for the full timeout. """ - if flush_event is None: - return - with contextlib.suppress(Exception): - flush_event.set() + if flush_event is not None: + with contextlib.suppress(Exception): + flush_event.set() def _reset_segment_state(self, *, preserve_no_edit: bool = False) -> None: if preserve_no_edit and self._message_id == "__no_edit__": @@ -566,45 +482,27 @@ class GatewayStreamConsumer( finalized = self._clean_for_display(self._last_sent_text).strip() if finalized: self._delivered_segment_texts.append(finalized) - self._message_id = None - self._message_created_ts = None - self._accumulated = "" - self._stream_ledger = "" - self._last_sent_text = "" - self._fallback_final_send = False - self._fallback_prefix = "" - self._fallback_preserve_partial_messages = False - self._segment_preview_message_ids = set() - self._tool_progress_lines = [] - self._tool_progress_active = False - # A segment boundary means what we delivered was an interim preamble — - # clear the final flags so a premature setter can't fool the gateway. - # Safe: got_done returns before any reset; run.py reads these only - # after the consumer task exits. - self._final_response_sent = False - self._final_content_delivered = False - self._delivered_final_text = None - self._delivery_ambiguous = False - self._turn_split_delivery = False - # Telegram-shaped drafts: bump draft_id so the next segment animates as - # a fresh preview below the tool-progress bubbles instead of over the - # prior finalized draft. Stream-is-the-message adapters (relay Slack) - # keep ONE stream per turn — a bump there opened a new platform stream - # per tool boundary, leaving one frozen cursor-suffixed message per - # segment; the connector's suffix-delta logic appends segments instead. + # Also clears the final flags: a segment boundary means what we delivered was + # an interim preamble, so a premature setter can't fool the gateway. Safe: + # got_done returns before any reset; run.py reads these after the task exits. + self._reset_message_state() + # Telegram-shaped drafts: bump draft_id so the next segment animates as a + # fresh preview below the tool-progress bubbles. Stream-is-the-message + # adapters (relay Slack) keep ONE stream per turn — a bump there left one + # frozen cursor-suffixed message per segment. if self._use_draft_streaming and not self._stream_is_message(): - type(self)._draft_id_counter += 1 - self._draft_id = type(self)._draft_id_counter + self._bump_draft_id() + + def _bump_draft_id(self) -> None: + type(self)._draft_id_counter += 1 + self._draft_id = type(self)._draft_id_counter async def _handle_approval_boundary(self, boundary_future, cancelled_flag=None) -> None: """Serially process an interaction boundary dequeued by run(). - Finalize the current stream with accumulated text (stable message for - pre-prompt content), then route post-prompt output per - ``_boundary_reopen``. The stream is not kept open across approval - because the WeCom finalize ack only confirms server receipt: after a - long idle gap the client may stop tracking the stream, and a suppressed - final send would then leave the user with nothing. + The stream is never kept open across a prompt: the WeCom finalize ack only + confirms server receipt, and after a long idle gap the client may stop + tracking the stream. """ _reason = self._boundary_reason or "Approval" try: @@ -612,52 +510,48 @@ class GatewayStreamConsumer( if self._native_stream_opened: boundary_ok = await self._finalize_boundary_stream(_reason) if self._boundary_reopen: - # Clarify: keep native enabled; marking the stream closed makes - # the next post-prompt delta re-open a fresh one via the lazy - # re-seed in _send_or_edit. Do NOT set buffer_only — post-prompt - # output should stream. The gap between this INFO and the - # "Re-opened native stream" INFO is the typing-reappear latency. + # Clarify: keep native enabled (NOT buffer_only); the closed stream + # makes the next post-prompt delta re-open via the lazy re-seed. The + # gap to the "Re-opened native stream" INFO is the typing latency. self._close_native_state() self._awaiting_reopen_after_boundary = True - self._reset_segment_state() + else: + # Approval: post-approval output goes via one send() at got_done. + self._degrade_native_to_buffered_send() + self._reset_segment_state() + if self._boundary_reopen: logger.info( "[latency] Clarify boundary finalized, awaiting first " "post-answer delta to re-seed (chat=%s, turn=%s)", self.chat_id, self._turn_id, ) - else: - # Approval: post-approval output goes via one send() at got_done. - self._degrade_native_to_buffered_send() - self._reset_segment_state() except Exception as e: logger.warning("%s boundary processing failed: %s", _reason, e) boundary_ok = False finally: - if isinstance(boundary_future, (asyncio.Future, concurrent.futures.Future)): - with contextlib.suppress(Exception): - if not boundary_future.done(): - boundary_future.set_result(boundary_ok) + with contextlib.suppress(Exception): + if ( + isinstance(boundary_future, (asyncio.Future, concurrent.futures.Future)) + and not boundary_future.done() + ): + boundary_future.set_result(boundary_ok) async def _finalize_boundary_stream(self, _reason: str) -> bool: """Close the open native stream at a boundary; send() the pre-prompt text if that fails. - Returns False only when both finalize and the fallback send failed - (pre-prompt text may not have been delivered). + Returns False only when both finalize and the fallback send failed. """ finalize_text = self._accumulated or self._boundary_placeholder - finalize_ok = False try: - finalize_ok = bool(await self._send_frame(finalize_text, finalize=True)) + if await self._send_frame(finalize_text, finalize=True): + logger.debug( + "%s boundary: finalized stream (chat=%s, turn=%s)", + _reason, self.chat_id, self._turn_id, + ) + return True except Exception as e: logger.warning("%s boundary: finalize failed: %s", _reason, e) - if finalize_ok: - logger.debug( - "%s boundary: finalized stream (chat=%s, turn=%s)", - _reason, self.chat_id, self._turn_id, - ) - return True - # Typing bubble may still show partial content; deliver pre-prompt - # text via send() so the user at least sees it. + # Typing bubble may still show partial content; deliver via send(). logger.warning( "%s boundary: finalize not confirmed, " "falling back to send() for pre-prompt text (chat=%s)", @@ -678,11 +572,10 @@ class GatewayStreamConsumer( return fallback_ok def on_delta(self, text: str) -> None: - """Thread-safe callback — called from the agent's worker thread. + """Thread-safe callback from the agent's worker thread. - When *text* is ``None``, signals a tool boundary: the current message - is finalized and subsequent text will be sent as a new message so it - appears below any tool-progress messages the gateway sent in between. + ``None`` signals a tool boundary: the current message is finalized and + subsequent text goes out as a new message below any tool-progress messages. """ if text: self._queue.put(text) @@ -693,9 +586,8 @@ class GatewayStreamConsumer( """Signal stream completion. ``final_text`` is the AUTHORITATIVE completed final_response (incl. - post-stream augmentation the accumulator never saw); the drain loop - adopts it as the finalize payload so no corrective send is needed. - Interrupt/error paths that can't know the final call ``finish()`` bare. + post-stream augmentation the accumulator never saw); the drain loop adopts + it as the finalize payload. Interrupt/error paths call ``finish()`` bare. """ if final_text is not None: self._queue.put((_FINAL_TEXT, final_text)) @@ -713,8 +605,8 @@ class GatewayStreamConsumer( return tick = self._drain_queue() - # Boundary produces its own finalize and resets state, so it - # must run before got_done/segment_break processing. + # Boundary produces its own finalize and resets state, so it must + # run before got_done/segment_break processing. if tick.approval_boundary is not None: await self._handle_approval_boundary(*tick.approval_boundary) continue @@ -724,12 +616,10 @@ class GatewayStreamConsumer( if tick.got_done: self._flush_think_buffer() - # A bare intentional-silence marker (NO_REPLY / [SILENT]): - # the gateway's whole-response filter runs too late for a - # streamed preview, so retract it here instead of finalizing. - if _is_intentional_silence_response( - self._clean_for_display(self._accumulated) - ): + # A bare intentional-silence marker (NO_REPLY / [SILENT]): the + # gateway's whole-response filter runs too late for a streamed + # preview, so retract it here instead of finalizing. + if _is_intentional_silence_response(self._clean_for_display(self._accumulated)): await self._suppress_silence_marker() return @@ -737,15 +627,14 @@ class GatewayStreamConsumer( self._accumulated or (self._use_native_streaming and self._tool_progress_active) ): - # Overflow split. Native streaming bypasses this: the - # adapter truncates against the stream protocol's own limit. + # Overflow split. Native streaming bypasses this: the adapter + # truncates against the stream protocol's own limit. if ( not self._use_native_streaming and self._len_fn(self._accumulated) > self._safe_limit and self._message_id is None ): - verdict = await self._split_first_send(tick) - if verdict == "return": + if await self._split_first_send(tick) == "return": return continue await self._seal_overflow_heads() @@ -777,54 +666,38 @@ class GatewayStreamConsumer( # ── run() collaborators ───────────────────────────────────────────── def _resolve_length_budget(self) -> "tuple[Callable[[str], int], int]": - """Per-chat length function + overflow budget. + """Per-chat length function (relay adapters differ per chat, e.g. utf16) + budget. - A relay adapter fronting N platforms has different caps per chat, in - the platform's unit (e.g. utf16 for Telegram). isinstance gate: - MagicMock auto-attributes aren't callables, so test doubles use len. + isinstance gate: MagicMock auto-attributes aren't callables; test doubles use len. """ len_fn: "Callable[[str], int]" = ( self.adapter.message_len_fn_for_chat(self.chat_id) if isinstance(self.adapter, _BasePlatformAdapter) else len ) - raw_limit = self._raw_message_limit() - return len_fn, max(500, raw_limit - len_fn(self.cfg.cursor) - 100) + return len_fn, max(500, self._raw_message_limit() - len_fn(self.cfg.cursor) - 100) async def _start_transports(self) -> None: """Resolve native/draft transport; native wins (adapters declaring it can't edit). - Sends an empty seed frame so the user sees "typing" before the first - token; on seed failure fall back to the edit path (the gateway's - fallback send then handles it). Drafts and native streaming target - the same first-frame slot. + The empty seed frame shows "typing" before the first token; on failure → edit path. """ self._use_native_streaming = self._resolve_native_streaming() if self._use_native_streaming: - logger.debug( - "Stream consumer using native-stream transport (chat=%s)", - self.chat_id, - ) + logger.debug("Stream consumer using native-stream transport (chat=%s)", self.chat_id) try: - seed_ok = await self._send_seed_frame() - if seed_ok: - self._native_stream_opened = True + seed_ok = bool(await self._send_seed_frame()) except Exception: - logger.debug( - "Native streaming seed frame raised; disabling native", - exc_info=True, - ) + logger.debug("Native streaming seed frame raised; disabling native", exc_info=True) seed_ok = False - if not seed_ok: - self._use_native_streaming = False - - if self._use_native_streaming: - self._use_draft_streaming = False - return + if seed_ok: + self._native_stream_opened = True + self._use_draft_streaming = False + return + self._use_native_streaming = False self._use_draft_streaming = self._resolve_draft_streaming() if self._use_draft_streaming: - type(self)._draft_id_counter += 1 - self._draft_id = type(self)._draft_id_counter + self._bump_draft_id() logger.debug( "Stream consumer using native-draft transport (chat=%s draft_id=%s)", self.chat_id, self._draft_id, @@ -833,9 +706,8 @@ class GatewayStreamConsumer( def _drain_queue(self) -> "_Tick": """Drain everything queued so far into one tick. - Control sentinels stop the drain (they take effect this tick); - _FINAL_TEXT / _TOOL_PROGRESS / text deltas fold into state and keep - draining so simultaneous items batch. + Control sentinels stop the drain (they take effect this tick); _FINAL_TEXT / + _TOOL_PROGRESS / text deltas fold into state so simultaneous items batch. """ tick = _Tick() while True: @@ -843,26 +715,17 @@ class GatewayStreamConsumer( item = self._queue.get_nowait() except queue.Empty: return tick - if item is _DONE: - tick.got_done = True - return tick - if item is _NEW_SEGMENT: - tick.got_segment_break = True - return tick - if item is _REOPEN_SEED: - tick.got_reopen_seed = True - return tick - if not isinstance(item, tuple): - self._filter_and_accumulate(item) - continue - try: - handler = self._QUEUE_TUPLE_HANDLERS.get((item[0], len(item))) - except TypeError: # unhashable head: not one of ours - handler = None + for sentinel, flag in self._QUEUE_SENTINEL_FLAGS: + if item is sentinel: + setattr(tick, flag, True) + return tick + handler = None + if isinstance(item, tuple): + with contextlib.suppress(TypeError): # unhashable head: not one of ours + handler = self._QUEUE_TUPLE_HANDLERS.get((item[0], len(item))) if handler is None: self._filter_and_accumulate(item) - continue - if handler(self, tick, item): + elif handler(self, tick, item): return tick def _on_final_text(self, tick: "_Tick", item: tuple) -> bool: @@ -894,21 +757,15 @@ class GatewayStreamConsumer( """Adopt the authoritative final (see finish()) as the finalize content. Only if this consumer streamed something (a no-stream turn keeps the - gateway's normal final-send ownership). Split delivery: wholesale - adoption would repeat sealed heads inside the tail, but refusing - entirely makes the gateway resend the ENTIRE body+footer — if the - final strictly prefix-extends the ledger, append only the missing - suffix; non-prefix rewrites keep the full-resend fallback. + gateway's final-send ownership). Split delivery: wholesale adoption would + repeat sealed heads, refusing makes the gateway resend the ENTIRE body — so + append only the suffix when the final strictly prefix-extends the ledger. """ - streamed_something = bool( - self._accumulated or self._message_id or self._last_sent_text - ) - if not streamed_something: + if not (self._accumulated or self._message_id or self._last_sent_text): return if not self._turn_split_delivery: final_payload = self._clean_for_display(final_raw) - visible = self._clean_for_display(self._accumulated) - if final_payload and final_payload != visible: + if final_payload and final_payload != self._clean_for_display(self._accumulated): self._accumulated = final_raw self._stream_ledger = final_raw return @@ -918,19 +775,12 @@ class GatewayStreamConsumer( self._stream_ledger = final_raw async def _eager_reopen_seed(self) -> None: - """Eager re-seed after a clarify answer. + """Eager re-seed after a clarify answer (gate re-checked: state may have advanced). - Re-check the gate (state may have advanced between put and dequeue). - Trade-off: WeCom's ~6-minute stream limit (errcode 846608, counted - from the FIRST frame) now starts at the reply instant rather than the - first post-answer delta; if it expires, send_stream_frame returns - False and we degrade to send() — only animation is lost. + Trade-off: WeCom's ~6-minute stream limit (errcode 846608, from the FIRST + frame) now starts at the reply instant; on expiry we degrade to send(). """ - if not ( - self._use_native_streaming - and self._awaiting_reopen_after_boundary - and not self._native_stream_opened - ): + if not self._reopen_seed_pending(): return try: seed_ok = await self._send_seed_frame() @@ -948,90 +798,75 @@ class GatewayStreamConsumer( self._turn_id, ) else: - # Degrade to a single buffered send() (one bubble, not per-tick - # fragments), like the approval path. + # Degrade to a single buffered send(), like the approval path. self._degrade_native_to_buffered_send() def _should_edit(self, tick: "_Tick") -> bool: """Decide whether this tick flushes an edit/frame.""" - should_edit = tick.got_done or tick.got_segment_break or tick.commentary_text is not None - if not self.cfg.buffer_only: - if self._use_native_streaming: - # No platform edit-rate limit: push every delta immediately. - should_edit = should_edit or bool(self._accumulated) or self._tool_progress_active - else: - elapsed = time.monotonic() - self._last_edit_time - should_edit = should_edit or ( - (elapsed >= self._current_edit_interval and self._accumulated) - # buffer_threshold is a codepoint debounce heuristic, - # not a platform-limit check (_len_fn is for overflow). - or len(self._accumulated) >= self.cfg.buffer_threshold - ) - # Defer mid-stream edits while the buffer could still resolve to a - # silence marker ("NO"→"NO_REPLY") so it never flashes on screen; - # got_done always resolves the buffer, so marker-like prose is never lost. - if ( - should_edit - and tick.is_interim - and _is_partial_silence_marker(self._clean_for_display(self._accumulated)) - ): + if not tick.is_interim: + return True + if self.cfg.buffer_only: return False - return should_edit + if self._use_native_streaming: + # No platform edit-rate limit: push every delta immediately. + should_edit = bool(self._accumulated) or self._tool_progress_active + else: + elapsed = time.monotonic() - self._last_edit_time + # buffer_threshold is a codepoint debounce heuristic, not a + # platform-limit check (_len_fn is for overflow). + should_edit = bool( + (elapsed >= self._current_edit_interval and self._accumulated) + or len(self._accumulated) >= self.cfg.buffer_threshold + ) + # Defer mid-stream edits while the buffer could still resolve to a silence + # marker ("NO"→"NO_REPLY"); got_done always resolves the buffer. + return should_edit and not _is_partial_silence_marker( + self._clean_for_display(self._accumulated) + ) async def _split_first_send(self, tick: "_Tick") -> str: """No message to edit yet and the buffer overflows: seal only the head chunks. - The tail stays in _accumulated so it becomes the active preview later - deltas edit in place. Returns "return" (turn finished) or "continue". + The tail stays in _accumulated so it becomes the active preview later deltas + edit in place. Returns "return" (turn finished) or "continue". """ chunks = self._truncate_for_stream(self._accumulated, self._safe_limit, self._len_fn) if len(chunks) <= 1: # Malformed/legacy adapter result must still be splittable. chunks = self._split_text_chunks(self._accumulated, self._safe_limit, self._len_fn) - chunks_delivered = False reply_to = self._initial_reply_to_id - all_heads_delivered = len(chunks) > 1 + heads_delivered = len(chunks) > 1 for chunk in chunks[:-1]: new_id = await self._send_new_chunk(chunk, reply_to, final=tick.got_done) if new_id is None or new_id == reply_to: - # Keep the full text intact for the gateway fallback. - all_heads_delivered = False - chunks_delivered = False + heads_delivered = False # keep the full text intact for the gateway fallback break - chunks_delivered = True reply_to = new_id - if all_heads_delivered: + if heads_delivered: self._accumulated = chunks[-1] - # Heads are sealed (or a later head failed): never edit a sealed - # message with the unsplit payload — the tail is sent fresh, or the - # fallback path retries. + # Flag BEFORE the tail send: fresh-final replaces every tracked preview + # with one message, which is only valid while the active message holds + # the whole answer — deleting sealed heads drops delivered text. + self._turn_split_delivery = True + # Heads are sealed (or a later head failed): never edit a sealed message with + # the unsplit payload — the tail is sent fresh, or the fallback path retries. self._message_id = None self._message_created_ts = None self._last_sent_text = "" - - if chunks_delivered: - # Flag BEFORE the tail send: fresh-final replaces every tracked - # preview with one message, which is only valid while the active - # message holds the whole answer — deleting sealed heads drops - # delivered text. - self._turn_split_delivery = True - self._last_edit_time = time.monotonic() if tick.got_done: tail_delivered = True if self._accumulated: tail_delivered = await self._send_or_edit(self._accumulated, finalize=True) - # ``_already_sent`` may be True from prior progress/fallback state - # — only heads + tail count. - self._final_response_sent = chunks_delivered and tail_delivered + # ``_already_sent`` may be True from prior state — only heads + tail count. + self._final_response_sent = heads_delivered and tail_delivered if self._final_response_sent: self._final_content_delivered = True self._turn_split_delivery = True self._record_turn_final_payload(self._accumulated) return "return" if tick.got_segment_break: - self._message_id = None self._fallback_final_send = False self._fallback_prefix = "" if not self._accumulated: @@ -1053,14 +888,12 @@ class GatewayStreamConsumer( if split_at < cp_budget // 2: split_at = cp_budget chunk = self._accumulated[:split_at] - # finalize=True: this sealed chunk is never edited again, so it - # needs its rich-text pass now or it renders raw. - # is_turn_final=False: a split head is not the answer, so - # fresh-final must not mark the turn delivered on it. + # finalize=True: the sealed chunk is never edited again, so it needs its + # rich-text pass now. is_turn_final=False: a split head is not the + # answer, so fresh-final must not mark the turn delivered on it. ok = await self._send_or_edit(chunk, finalize=True, is_turn_final=False) if self._fallback_final_send or not ok: - # Keep the full text intact for the fallback final send. - break + break # keep the full text intact for the fallback final send self._accumulated = self._accumulated[split_at:].lstrip("\n") self._message_id = None self._last_sent_text = "" @@ -1077,39 +910,30 @@ class GatewayStreamConsumer( else: display_text += self.cfg.cursor - # got_done update going out as a FRESH send via the draft transport - # (drafts have no message id) already carries finalize=True — unlike an - # EDIT during draft streaming, which still needs the explicit finalize - # pass for REQUIRES_EDIT_FINALIZE adapters. + # A got_done FRESH send via the draft transport already carries finalize=True, + # unlike an EDIT, which REQUIRES_EDIT_FINALIZE adapters still need a pass for. tick.draft_final_fresh_send = ( tick.got_done and self._use_draft_streaming and self._message_id is None ) - # Segment break finalizes so platforms needing explicit closure - # (DingTalk AI Cards) don't leave the previous segment stuck loading; - # got_done has its own finalize in _finalize_turn. + # Segment break finalizes so platforms needing explicit closure (DingTalk AI + # Cards) don't leave the segment stuck loading; it closes a preamble, not the + # answer. tick.update_visible = await self._send_or_edit( display_text, finalize=(tick.got_done or tick.got_segment_break), - # A segment-break finalize closes a preamble, not the answer. is_turn_final=tick.got_done, ) self._last_edit_time = time.monotonic() - # Lines stay in _tool_progress_lines for the next compose; only new - # progress should trigger another should_edit. - if self._tool_progress_active: - self._tool_progress_active = False + # Lines stay in _tool_progress_lines for the next compose. + self._tool_progress_active = False async def _finalize_turn(self, tick: "_Tick") -> None: - """got_done: final edit without cursor, or one continuation send if edits failed mid-stream.""" + """got_done: final edit without cursor, or one continuation send if edits failed.""" if self._accumulated or self._message_id is not None or self._already_sent: await self._notify_before_finalize() - if ( - self._awaiting_reopen_after_boundary - and not self._native_stream_opened - and not self._accumulated - ): - # Lazy reopen, no post-prompt content: nothing is open on screen, - # so don't re-seed just to emit a lone "✅". + if self._reopen_seed_pending() and not self._accumulated: + # Lazy reopen, no post-prompt content: nothing is open on screen, so + # don't re-seed just to emit a lone "✅". logger.debug( "Clarify reopen boundary with no post-prompt content " "— skipping lone-placeholder finalize (turn=%s)", @@ -1121,23 +945,17 @@ class GatewayStreamConsumer( and not self._accumulated and not tick.update_visible ): - # Eager seed, no content: the typing bubble IS on screen and would - # hang forever — close it with an empty finalize (not a lone "✅"). - # Delivery flags untouched. - try: - await self._send_frame("", finalize=True) - except Exception as e: - logger.debug("Eager-seed empty finalize failed: %s", e) - self._close_native_state() - self._reopen_seeded_eagerly = False + # Eager seed, no content: the typing bubble IS on screen and would hang + # forever — close it with an empty finalize. Delivery flags untouched. + await self._close_empty_native_bubble("Eager-seed empty finalize failed: %s") logger.debug( "Eager reopen seed but no post-answer content — " "closed empty typing bubble (turn=%s)", self._turn_id, ) elif self._use_native_streaming: - # Native streams MUST close with finish=true even when empty - # (tool-only turns) — placeholder if needed. + # Native streams MUST close with finish=true even when empty (tool-only + # turns) — placeholder if needed. if not tick.update_visible: await self._finalize_edit(self._accumulated or "✅", record=False) else: @@ -1146,12 +964,11 @@ class GatewayStreamConsumer( await self._finalize_edit_path(tick) async def _finalize_edit_path(self, tick: "_Tick") -> None: - """Edit-transport finalize (the non-native got_done branches).""" + """Edit-transport finalize (the non-native got_done branches, in priority order).""" if self._fallback_final_send: await self._send_fallback_final(self._accumulated) elif self._final_response_sent: - # Fresh-final already delivered above; a second finalize would - # duplicate / re-delete. + # Fresh-final already delivered; a second finalize would duplicate. self._final_content_delivered = True self._record_turn_final_payload(self._accumulated) elif tick.update_visible and ( @@ -1159,29 +976,26 @@ class GatewayStreamConsumer( or self._last_edit_overflowed or tick.draft_final_fresh_send ): - # The update above already delivered the final. A second finalize - # would re-edit an already-final message (Telegram: sendRichMessage - # followed by editMessageText falls back to the legacy formatter), - # or overflow-split again into an adopted continuation, - # duplicating chunks. + # The update already delivered the final. A second finalize would + # re-edit an already-final message (Telegram: sendRichMessage followed by + # editMessageText falls back to the legacy formatter) or overflow-split + # again into an adopted continuation, duplicating chunks. self._mark_skip_redundant_finalize() elif self._message_id: # No visible update this tick, or the adapter needs explicit finalize=True. - ok = await self._finalize_edit(self._accumulated) - if not ok and self._fallback_final_send: - # This edit may have exhausted flood strikes and promoted - # fallback mode; send only the unsent tail now so the gateway - # doesn't duplicate the visible prefix. + # The edit may exhaust flood strikes and promote fallback mode: then send + # only the unsent tail. + if not await self._finalize_edit(self._accumulated) and self._fallback_final_send: await self._send_fallback_final(self._accumulated) elif not self._already_sent: # Retry after the finalize tick failed. finalize=True keeps - # stream-is-the-message adapters out of the draft-frame branch, - # whose dedupe against the last UNSEALED frame would report - # success with no transport call (silent loss). + # stream-is-the-message adapters out of the draft-frame branch, whose + # dedupe against the last UNSEALED frame would report success with no + # transport call (silent loss). await self._finalize_edit(self._accumulated) async def _finalize_edit(self, text: str, *, record: bool = True) -> bool: - """finalize=True send_or_edit; on success mark the turn delivered (and record the payload).""" + """finalize=True send_or_edit; on success mark the turn delivered (+ record payload).""" self._final_response_sent = await self._send_or_edit(text, finalize=True) if self._final_response_sent: self._mark_final_delivered(record=text if record else None) @@ -1194,34 +1008,32 @@ class GatewayStreamConsumer( ) or self._use_native_streaming async def _deliver_commentary(self, commentary_text: str) -> None: - """Cumulative transports post commentary as its own message and the - stream continues — resetting _accumulated would break the append-only - invariant / lose pre-commentary text.""" - if self._cumulative_transport(): - await self._send_commentary(commentary_text) - self._last_edit_time = time.monotonic() - else: + """Post commentary as its own message. + + Cumulative transports keep the stream going — resetting _accumulated would + break the append-only invariant / lose pre-commentary text. + """ + cumulative = self._cumulative_transport() + if not cumulative: self._reset_segment_state() - await self._send_commentary(commentary_text) - self._last_edit_time = time.monotonic() + await self._send_commentary(commentary_text) + self._last_edit_time = time.monotonic() + if not cumulative: self._reset_segment_state() async def _end_segment(self, tick: "_Tick") -> None: - """Tool boundary: edit-based transports reset so the next chunk is a fresh message below tool progress. + """Tool boundary: edit-based transports reset so the next chunk is a fresh message. - Cumulative transports must NOT reset — clearing _accumulated makes the - next frame a non-prefix snapshot and the connector re-appends the - whole answer. preserve_no_edit: "__no_edit__" (platform never - returned a real id — Signal, github_comment webhook) must keep its - sentinel or every tool boundary posts a new message (155 comments on - one PR); the continuation goes out once via _send_fallback_final. A - real id from a flood-failed edit still resets as intended. + Cumulative transports must NOT reset — clearing _accumulated makes the next + frame a non-prefix snapshot and the connector re-appends the whole answer. + preserve_no_edit: "__no_edit__" (platform never returned a real id — Signal, + github_comment webhook) must keep its sentinel or every tool boundary posts a + new message; the continuation goes out once via _send_fallback_final. """ if self._cumulative_transport(): return - # If the segment-break edit didn't land (flood control / fallback - # mode), _accumulated holds unseen pre-boundary text — flush it - # before the reset wipes it. + # If the segment-break edit didn't land (flood control / fallback mode), + # _accumulated holds unseen pre-boundary text — flush it before the reset. if ( self._accumulated and not tick.update_visible @@ -1234,11 +1046,9 @@ class GatewayStreamConsumer( async def _on_cancelled(self) -> None: """Best-effort final edit on task cancel. - finalize=True so REQUIRES_EDIT_FINALIZE platforms apply formatting - (the flags set here suppress the gateway's formatted re-send); - is_turn_final=False because this handler owns the flags, not - _try_fresh_final. Only a successful best-effort edit confirms - delivery — a partial send (already_sent) may be just "Let me + finalize=True so REQUIRES_EDIT_FINALIZE platforms apply formatting; + is_turn_final=False because this handler owns the flags. Only a successful + best-effort edit confirms delivery — a partial send may be just "Let me search…", not the answer. """ best_effort_ok = False @@ -1250,10 +1060,8 @@ class GatewayStreamConsumer( ) ) elif self._message_id is None: - # Draft path keeps _message_id=None, so the edit above never runs - # for it; seal in place (else the stream stays visibly live and - # the adapter keeps armed interception state for the next turn). - # Sets no delivery flags. + # Draft path keeps _message_id=None; seal in place (else the stream stays + # visibly live and the adapter keeps armed interception state). await self._abandon_native_stream() if best_effort_ok and not self._final_response_sent: self._mark_final_delivered(record=self._accumulated) @@ -1261,15 +1069,11 @@ class GatewayStreamConsumer( def _wake_flush_waiters(self) -> None: """Wake still-queued _FLUSH waiters so a consumer dying mid-flush doesn't stall flush_pending_sync() for its full timeout.""" - try: + with contextlib.suppress(Exception): while True: item = self._queue.get_nowait() if isinstance(item, tuple) and len(item) == 2 and item[0] is _FLUSH: self._signal_flush(item[1]) - except queue.Empty: - pass - except Exception: - pass # Tuple-shaped queue items keyed on (sentinel, arity); handler returns True # to stop draining. Order-insensitive: each sentinel is a distinct object. @@ -1280,9 +1084,14 @@ class GatewayStreamConsumer( (_FLUSH, 2): _on_flush, (_TOOL_PROGRESS, 2): _on_tool_progress, } + # Bare control sentinels -> the _Tick flag they set (each ends the drain). + _QUEUE_SENTINEL_FLAGS = ( + (_DONE, "got_done"), + (_NEW_SEGMENT, "got_segment_break"), + (_REOPEN_SEED, "got_reopen_seed"), + ) @staticmethod def _clean_for_display(text: str) -> str: - """Hide MEDIA: / [[audio_as_voice]] directives; media is delivered after the stream.""" + """Hide MEDIA: / [[audio_as_voice]] directives; media is delivered post-stream.""" return _BasePlatformAdapter.strip_media_directives_for_display(text) - diff --git a/gateway/stream_consumer_fallback.py b/gateway/stream_consumer_fallback.py index 556effa10f..21f9ea2490 100644 --- a/gateway/stream_consumer_fallback.py +++ b/gateway/stream_consumer_fallback.py @@ -4,6 +4,7 @@ stop working, chunking, cursor cleanup, commentary and silence retraction.""" from __future__ import annotations import asyncio +import contextlib import logging from typing import Any, Callable, Optional @@ -32,21 +33,17 @@ class StreamFallbackMixin: chat_id=self.chat_id, content=text, reply_to=reply_to_id, - metadata=self._metadata_for_send( - final=final, - expect_edits=not final, - ), + metadata=self._metadata_for_send(final=final, expect_edits=not final), ) - if result.success and result.message_id: - self._message_id = str(result.message_id) - self._track_preview_ids_from_result(result) - self._already_sent = True - self._last_sent_text = text - self._notify_new_message() - return str(result.message_id) - else: + if not (result.success and result.message_id): self._edit_supported = False return reply_to_id + self._message_id = str(result.message_id) + self._track_preview_ids_from_result(result) + self._already_sent = True + self._last_sent_text = text + self._notify_new_message() + return str(result.message_id) except Exception as e: logger.error("Stream send chunk error: %s", e) return reply_to_id @@ -67,26 +64,17 @@ class StreamFallbackMixin: @staticmethod def _split_text_chunks( - text: str, - limit: int, - len_fn: "Callable[[str], int]" = len, + text: str, limit: int, len_fn: "Callable[[str], int]" = len, ) -> list[str]: """Split text for fallback sends: newline-preferred, fence-balanced across chunks.""" from gateway.platforms.helpers import split_text_fence_aware return split_text_fence_aware( - text, - limit, - len_fn, - prefer_paragraphs=False, - balance_fences=True, + text, limit, len_fn, prefer_paragraphs=False, balance_fences=True, ) def _truncate_for_stream( - self, - text: str, - limit: int, - len_fn: "Callable[[str], int]", + self, text: str, limit: int, len_fn: "Callable[[str], int]", ) -> list[str]: """Split via the adapter's canonical truncate_message (platform-specific rules). @@ -95,14 +83,11 @@ class StreamFallbackMixin: truncate = getattr(self.adapter, "truncate_message", None) if not callable(truncate): return self._split_text_chunks(text, limit, len_fn) - if isinstance(self.adapter, _BasePlatformAdapter): chunks = truncate(text, limit, len_fn=len_fn) else: chunks = truncate(text, limit) - if not isinstance(chunks, (list, tuple)) or not all( - isinstance(chunk, str) for chunk in chunks - ): + if not isinstance(chunks, (list, tuple)) or not all(isinstance(c, str) for c in chunks): return self._split_text_chunks(text, limit, len_fn) return list(chunks) @@ -111,10 +96,9 @@ class StreamFallbackMixin: Retries each chunk once on flood-control failures with a short delay. """ - final_text = self._clean_for_display(text) - # Balance fences BEFORE computing the continuation so the closing - # fence reaches the user even when only the tail is delivered. - final_text = ensure_closed_code_fences(final_text) + # Balance fences BEFORE computing the continuation so the closing fence + # reaches the user even when only the tail is delivered. + final_text = ensure_closed_code_fences(self._clean_for_display(text)) continuation = self._continuation_text(final_text) self._fallback_final_send = False if not continuation.strip(): @@ -123,8 +107,7 @@ class StreamFallbackMixin: return _len_fn, raw_limit = self._fallback_len_budget() - safe_limit = max(500, raw_limit - 100) - chunks = self._split_text_chunks(continuation, safe_limit, len_fn=_len_fn) + chunks = self._split_text_chunks(continuation, max(500, raw_limit - 100), len_fn=_len_fn) stale_message_id = self._message_id # partial message to clean up last_message_id: Optional[str] = None @@ -135,19 +118,13 @@ class StreamFallbackMixin: content=chunk, retry_log="Flood control on fallback send, retrying in %.1fs", ) if not result or not result.success: - if sent_any_chunk: - # Partial continuation landed: do NOT set _final_response_sent - # (gateway must still deliver the full answer); _already_sent - # only prevents a duplicate of the partial. - self._already_sent = True - self._message_id = last_message_id - self._last_sent_text = last_successful_chunk - self._fallback_prefix = "" - return - # Nothing landed — let the gateway final send try once more. - self._already_sent = False - self._message_id = None - self._last_sent_text = "" + # Partial continuation landed: do NOT set _final_response_sent (the + # gateway must still deliver the full answer); _already_sent only + # prevents a duplicate of the partial. Nothing landed: let the + # gateway final send try once more. + self._already_sent = sent_any_chunk + self._message_id = last_message_id + self._last_sent_text = last_successful_chunk self._fallback_prefix = "" return sent_any_chunk = True @@ -155,9 +132,9 @@ class StreamFallbackMixin: last_message_id = result.message_id or last_message_id self._notify_new_message() - # Best-effort delete of the frozen partial — ONLY when the FULL final - # was re-sent. If only the missing tail went out, the partial IS the - # head of the answer ("sent only the second half" symptom). + # Best-effort delete of the frozen partial — ONLY when the FULL final was + # re-sent. If only the missing tail went out, the partial IS the head of + # the answer ("sent only the second half" symptom). if ( stale_message_id and stale_message_id != last_message_id @@ -179,14 +156,15 @@ class StreamFallbackMixin: async def _fallback_when_nothing_unseen(self, final_text: str) -> Optional[str]: """Fallback entered but the visible prefix already covers ``final_text``. - Returns the continuation to send (the whole final when the prefix is - from a *previous* segment), or None when the turn is settled here. + Returns the continuation to send (the whole final when the prefix is from a + *previous* segment) or None when the turn is settled here. """ - # Telegram clients can lose (part of) a streamed preview after a - # failed final edit, so opt-in adapters commit a fresh final send. + visible = self._visible_prefix() + # Telegram clients can lose (part of) a streamed preview after a failed + # final edit, so opt-in adapters commit a fresh final send. if ( final_text.strip() - and final_text == self._visible_prefix() + and final_text == visible and getattr(self.adapter, "RESEND_FINAL_ON_EMPTY_STREAM_FALLBACK", False) is True ): delivery = await self._send_empty_fallback_final(final_text) @@ -196,13 +174,10 @@ class StreamFallbackMixin: self._fallback_prefix = "" self._fallback_preserve_partial_messages = False if delivery in {"ambiguous", "preview"}: - # Timeout: Telegram may have accepted the send. Flood - # rejection: the complete ACKed preview is authoritative. - # Keep duplicate suppression in both cases. + # Timeout: Telegram may have accepted the send. Flood rejection: the + # complete ACKed preview is authoritative. Keep dup suppression. self._final_content_delivered = True if delivery == "preview": - # Preview already shows the full final (checked above); - # record it so the gateway doesn't re-send next to it. self._record_turn_final_payload(final_text) else: self._delivery_ambiguous = True @@ -211,12 +186,11 @@ class StreamFallbackMixin: self._final_response_sent = False self._final_content_delivered = False return None - # The prefix may be from a *previous* segment (before a tool - # boundary), wrongly reading as "already shown" — send final_text as-is. - if final_text.strip() and final_text != self._visible_prefix(): + # The prefix may be from a *previous* segment (before a tool boundary), + # wrongly reading as "already shown" — send final_text as-is. + if final_text.strip() and final_text != visible: return final_text - # Best-effort strip of a cursor left stuck by the edit failure that - # entered fallback mode. + # Best-effort strip of a cursor left stuck by the edit failure. if ( self._message_id and self._last_sent_text @@ -224,12 +198,10 @@ class StreamFallbackMixin: and self._last_sent_text.endswith(self.cfg.cursor) ): clean_text = self._last_sent_text[:-len(self.cfg.cursor)] - try: + with contextlib.suppress(Exception): result = await self._edit_message(message_id=self._message_id, content=clean_text) if result.success: self._last_sent_text = clean_text - except Exception: - pass self._already_sent = True # Recorder substitutes the full ledger on a split turn. self._mark_final_delivered(record=final_text) @@ -241,8 +213,7 @@ class StreamFallbackMixin: _len_fn: "Callable[[str], int]" = len if isinstance(self.adapter, _BasePlatformAdapter): _len_fn = self.adapter.message_len_fn - # Per-chat cap/unit (relay adapter fronting N platforms). - try: + try: # per-chat cap/unit (relay adapter fronting N platforms) raw_limit = self.adapter.max_message_length_for_chat(self.chat_id) _len_fn = self.adapter.message_len_fn_for_chat(self.chat_id) except Exception as e: @@ -255,32 +226,28 @@ class StreamFallbackMixin: Exceptions propagate (callers decide whether a raise means "ambiguous"). Returns the last SendResult (success or not). """ + kwargs = dict( + chat_id=self.chat_id, content=content, metadata=self._metadata_for_send(final=True), + ) + if reply_to is not None: + kwargs["reply_to"] = reply_to result = None for attempt in range(2): - kwargs = dict( - chat_id=self.chat_id, - content=content, - metadata=self._metadata_for_send(final=True), - ) - if reply_to is not None: - kwargs["reply_to"] = reply_to result = await self.adapter.send(**kwargs) if getattr(result, "success", False): break retry_delay = self._fallback_flood_retry_delay(result) - if attempt == 0 and retry_delay is not None: - logger.debug(retry_log, retry_delay) - await asyncio.sleep(retry_delay) - else: + if attempt or retry_delay is None: break # non-flood error, long flood wait, or second failure + logger.debug(retry_log, retry_delay) + await asyncio.sleep(retry_delay) return result async def _send_empty_fallback_final(self, final_text: str) -> str: """Commit a completed answer after Telegram finalization fails. - Returns "delivered", "failed" (gateway may retry), "ambiguous" (a - timeout may have reached the platform), or "preview" (flood control - leaves the complete streamed preview authoritative). + Returns "delivered", "failed" (gateway may retry), "ambiguous" (a timeout may + have landed) or "preview" (flood control; the complete preview is authoritative). """ # Segment-scoped only: never delete an earlier finalized preamble. stale_ids = self._stale_preview_ids(segment_only=True) @@ -299,8 +266,8 @@ class StreamFallbackMixin: return "ambiguous" if self._send_failure_may_have_delivered(result) else "failed" new_message_id = getattr(result, "message_id", None) - # Telegram reports delete failure by returning False; the flood window - # that broke the finalize can reject this too — one bounded retry. + # Telegram reports delete failure by returning False; the flood window that + # broke the finalize can reject this too — one bounded retry. await self._delete_previews( stale_ids, skip=new_message_id, label="Empty fallback", retry_on_false=True, ) @@ -308,12 +275,10 @@ class StreamFallbackMixin: self._message_id = new_message_id or "__no_edit__" self._already_sent = True self._mark_final_delivered() - # Record VERBATIM, not via _record_turn_final_payload: the sealed - # previews were just deleted, so the ledger (which still holds sealed - # heads) would claim delivery for text this path removed. - self._delivered_final_text = ensure_closed_code_fences( - self._clean_for_display(final_text or "") - ).strip() + # Record VERBATIM, not via _record_turn_final_payload: the sealed previews + # were just deleted, so the ledger (still holding sealed heads) would claim + # delivery for text this path removed. + self._delivered_final_text = self._display_payload(final_text) self._last_sent_text = final_text self._fallback_prefix = "" self._fallback_preserve_partial_messages = False @@ -347,8 +312,7 @@ class StreamFallbackMixin: def _is_flood_error(self, result) -> bool: """Check if a SendResult failure is due to flood control / rate limiting.""" - err = getattr(result, "error", "") or "" - err_lower = err.lower() + err_lower = (getattr(result, "error", "") or "").lower() return "flood" in err_lower or "retry after" in err_lower or "rate" in err_lower async def _flush_segment_tail_on_edit_failure(self) -> None: @@ -369,11 +333,7 @@ class StreamFallbackMixin: # Interim: must never seal a native stream (see _send_commentary). _md = dict(self.metadata) if self.metadata else {} _md["_interim_send"] = True - result = await self.adapter.send( - chat_id=self.chat_id, - content=tail, - metadata=_md, - ) + result = await self.adapter.send(chat_id=self.chat_id, content=tail, metadata=_md) if result.success: self._already_sent = True except Exception as e: @@ -384,17 +344,12 @@ class StreamFallbackMixin: if not self._message_id or self._message_id == "__no_edit__": return prefix = self._visible_prefix() - if not prefix or not prefix.strip(): + if not prefix.strip(): return - try: - result = await self._edit_message( - message_id=self._message_id, - content=prefix, - ) + with contextlib.suppress(Exception): # never block the fallback path + result = await self._edit_message(message_id=self._message_id, content=prefix) if getattr(result, "success", False): self._last_sent_text = prefix - except Exception: - pass # best-effort — don't let this block the fallback path async def _send_commentary(self, text: str) -> bool: """Send a completed interim assistant commentary message.""" @@ -402,9 +357,9 @@ class StreamFallbackMixin: if not text.strip(): return False try: - # Interim: a stream-is-the-message adapter's seal-interception must - # not turn this into draft(final=true), which would seal the live - # stream with interim text and orphan the true final. + # Interim: a stream-is-the-message adapter's seal-interception must not + # turn this into draft(final=true), which would seal the live stream + # with interim text and orphan the true final. _md = self._metadata_for_send(final=False) or {} _md["_interim_send"] = True # reply_to only for reply-anchored threading; Discord/Telegram use @@ -418,8 +373,8 @@ class StreamFallbackMixin: reply_to=self._initial_reply_to_id if _needs_reply_anchor else None, metadata=_md, ) - # Do NOT set _already_sent: commentary is interim, and the flag - # would suppress the real final after multiple tool calls. + # Do NOT set _already_sent: commentary is interim, and the flag would + # suppress the real final after multiple tool calls. if result.success: self._notify_new_message() # Lets run.py confirm whether an interim send carried the final. @@ -454,36 +409,18 @@ class StreamFallbackMixin: async def _suppress_silence_marker(self) -> None: """Retract any streamed preview when the final reply is a bare silence marker. - Delivery flags and ``_already_sent`` are left False: nothing was - delivered, and the gateway's whole-response filter turns the marker - into "" so no fallback send happens either. + Flags stay False: nothing was delivered, and the gateway's whole-response + filter turns the marker into "" so no fallback send happens either. """ # A native-stream bubble isn't a deletable message — close an open one # (e.g. from an eager re-seed) with an empty finalize so it doesn't hang. if self._native_stream_opened: - try: - await self._send_frame("", finalize=True) - except Exception as e: - logger.debug( - "Silence-marker native stream close failed: %s", e, - ) - self._close_native_state() - self._reopen_seeded_eagerly = False + await self._close_empty_native_bubble("Silence-marker native stream close failed: %s") - stale_ids = self._stale_preview_ids() - await self._delete_previews(stale_ids, label="Silence-marker") + await self._delete_previews(self._stale_preview_ids(), label="Silence-marker") self._preview_message_ids = set() self._message_id = None - self._accumulated = "" - self._stream_ledger = "" - self._last_sent_text = "" + self._accumulated = self._stream_ledger = self._last_sent_text = "" self._already_sent = False - self._final_response_sent = False - self._final_content_delivered = False - self._delivered_final_text = None - self._delivery_ambiguous = False - self._turn_split_delivery = False - logger.info( - "Suppressed streamed intentional-silence marker (chat=%s)", - self.chat_id, - ) + self._clear_turn_final_flags() + logger.info("Suppressed streamed intentional-silence marker (chat=%s)", self.chat_id) diff --git a/gateway/stream_consumer_fences.py b/gateway/stream_consumer_fences.py index c0dde5db3a..6c7daa7d16 100644 --- a/gateway/stream_consumer_fences.py +++ b/gateway/stream_consumer_fences.py @@ -10,20 +10,16 @@ def escape_code_fences_for_display(text: str) -> str: Reasoning content that quotes code would otherwise break the outer fence. """ - if not isinstance(text, str) or "```" not in text: - return text - return text.replace("```", "\\`\\`\\`") + return text.replace("```", "\\`\\`\\`") if isinstance(text, str) else text def ensure_closed_code_fences(text: str) -> str: """Append a closing ``` and/or ` if the text has orphaned code markers. - Output truncated mid-code-block (token limit, finish_reason="length") would - otherwise render everything after the orphan as one code block / inline - span. Trade-off: a spurious close creates a brief empty span at the end, - far less harmful than the alternative. Odd ``` count → append a fence on - its own line; then, with complete ```…``` regions stripped, odd ` count → - append a backtick. + Output truncated mid-code-block (finish_reason="length") would otherwise render + everything after the orphan as one code block / inline span; a spurious close is + far less harmful. Odd ``` count → fence on its own line; then, with complete + ```…``` regions stripped, odd ` count → a backtick. """ if not isinstance(text, str) or not text: return text diff --git a/gateway/stream_consumer_think.py b/gateway/stream_consumer_think.py index 5ebe3542c5..baab0f8c7f 100644 --- a/gateway/stream_consumer_think.py +++ b/gateway/stream_consumer_think.py @@ -25,6 +25,34 @@ class StreamThinkFilterMixin: "", "", "", ) + def _at_block_boundary(self, buf: str, idx: int) -> bool: + """Tag at ``idx`` starts a block: start of text, or newline + optional whitespace. + + Prose that merely *mentions* a tag must not trigger (mirrors cli.py). + """ + acc_boundary = not self._accumulated or self._accumulated.endswith("\n") + if idx == 0: + return acc_boundary + preceding = buf[:idx] + last_nl = preceding.rfind("\n") + if last_nl == -1: + return acc_boundary and preceding.strip() == "" + return preceding[last_nl + 1:].strip() == "" + + def _earliest_open_tag(self, buf: str, lower_buf: str) -> "tuple[int, int]": + """(index, length) of the earliest block-boundary opening tag, or (-1, 0).""" + best_idx, best_len = -1, 0 + for tag in self._OPEN_THINK_TAGS: + tag_lower = tag.lower() + search_start = 0 + while (idx := lower_buf.find(tag_lower, search_start)) != -1: + if self._at_block_boundary(buf, idx): + if best_idx == -1 or idx < best_idx: + best_idx, best_len = idx, len(tag) + break # first boundary hit for this tag is enough + search_start = idx + 1 + return best_idx, best_len + def _filter_and_accumulate(self, text: str) -> None: """Append a delta to the buffer, discarding think blocks. @@ -38,14 +66,11 @@ class StreamThinkFilterMixin: # Case-insensitive: models emit , , … lower_buf = buf.lower() if self._in_think_block: - best_idx = -1 - best_len = 0 + best_idx, best_len = -1, 0 for tag in self._CLOSE_THINK_TAGS: idx = lower_buf.find(tag.lower()) if idx != -1 and (best_idx == -1 or idx < best_idx): - best_idx = idx - best_len = len(tag) - + best_idx, best_len = idx, len(tag) if best_len: self._in_think_block = False buf = buf[best_idx + best_len:] @@ -55,60 +80,24 @@ class StreamThinkFilterMixin: self._think_buffer = buf[-max_tag:] if len(buf) > max_tag else buf return else: - # Earliest opening tag at a block boundary (start of text, or - # newline + optional whitespace) — prose that merely *mentions* - # a tag must not trigger. - best_idx = -1 - best_len = 0 - for tag in self._OPEN_THINK_TAGS: - tag_lower = tag.lower() - search_start = 0 - while True: - idx = lower_buf.find(tag_lower, search_start) - if idx == -1: - break - # Block-boundary check (mirrors cli.py logic) - if idx == 0: - is_boundary = ( - not self._accumulated - or self._accumulated.endswith("\n") - ) - else: - preceding = buf[:idx] - last_nl = preceding.rfind("\n") - if last_nl == -1: - is_boundary = ( - (not self._accumulated - or self._accumulated.endswith("\n")) - and preceding.strip() == "" - ) - else: - is_boundary = preceding[last_nl + 1:].strip() == "" - - if is_boundary and (best_idx == -1 or idx < best_idx): - best_idx = idx - best_len = len(tag) - break # first boundary hit for this tag is enough - search_start = idx + 1 - + best_idx, best_len = self._earliest_open_tag(buf, lower_buf) if best_len: self._append_accumulated(buf[:best_idx]) self._in_think_block = True buf = buf[best_idx + best_len:] else: # Hold back a partial open tag at the tail. - held_back = 0 - for tag in self._OPEN_THINK_TAGS: - tag_lower = tag.lower() - for i in range(1, len(tag)): - if lower_buf.endswith(tag_lower[:i]) and i > held_back: - held_back = i + held_back = max( + (i for tag in self._OPEN_THINK_TAGS for i in range(1, len(tag)) + if lower_buf.endswith(tag.lower()[:i])), + default=0, + ) if held_back: self._append_accumulated(buf[:-held_back]) self._think_buffer = buf[-held_back:] else: - # An orphan (thinking-mode toggle dropped the - # open, or incomplete upstream stripping) is noise. + # An orphan (thinking-mode toggle dropped the open, or + # incomplete upstream stripping) is noise. self._append_accumulated(self._strip_orphan_close_tags(buf)) return @@ -129,9 +118,8 @@ class StreamThinkFilterMixin: if text_lower[i:i + 2] == " None: + """Best-effort empty finalize frame to close an open typing bubble, then mark closed.""" + try: + await self._send_frame("", finalize=True) + except Exception as e: + logger.debug(fail_log, e) + self._close_native_state() + self._reopen_seeded_eagerly = False + def _degrade_native_to_buffered_send(self) -> None: """Leave native mode; post-boundary output goes out as ONE send() at got_done. @@ -84,13 +83,8 @@ class StreamTransportMixin: self.cfg.buffer_only = True def _draft_metadata(self) -> dict | None: - """Draft-frame metadata. - - Every frame must carry the same reply_to_message_id the final send gets - from _metadata_for_send: the relay adapter keys draft/seal state on it, - else the final can't find the open stream (flat DMs have no thread - metadata and would key on the bare chat). - """ + """Draft-frame metadata: same reply_to_message_id as the final send, because the + relay adapter keys draft/seal state on it (flat DMs have no thread metadata).""" md = dict(self.metadata) if self.metadata else {} if self._initial_reply_to_id: md.setdefault("reply_to_message_id", self._initial_reply_to_id) @@ -101,14 +95,11 @@ class StreamTransportMixin: ``segment_only``: never delete an earlier finalized preamble. """ - if segment_only: - stale_ids = set(self._segment_preview_message_ids) - if self._message_id and self._message_id != "__no_edit__": - stale_ids.add(str(self._message_id)) - return stale_ids - stale_ids = set(self._preview_message_ids) + stale_ids = set( + self._segment_preview_message_ids if segment_only else self._preview_message_ids + ) if self._message_id and self._message_id != "__no_edit__": - stale_ids.add(self._message_id) + stale_ids.add(str(self._message_id) if segment_only else self._message_id) return stale_ids async def _delete_previews( @@ -133,9 +124,8 @@ class StreamTransportMixin: def _resolve_draft_streaming(self) -> bool: """Whether this run should use draft streaming per ``cfg.transport``. - "edit"/"off" → False. "draft"/"auto" → the adapter's - supports_draft_streaming probe (chat type, platform-version gates); - "draft" logs the downgrade when unsupported. + "edit"/"off" → False. "draft"/"auto" → the adapter's supports_draft_streaming + probe (chat type, platform-version gates); "draft" logs the downgrade. """ transport = (self.cfg.transport or "edit").lower() if transport in ("edit", "off"): @@ -143,32 +133,26 @@ class StreamTransportMixin: # MagicMock test adapters default to edit. if not isinstance(self.adapter, _BasePlatformAdapter): return False + probe_kwargs = dict(chat_type=self.cfg.chat_type or None, metadata=self.metadata) try: try: # Per-chat probe (relay adapters resolve through the CHAT's # descriptor); older adapters without the kwarg keep the legacy probe. supported = self.adapter.supports_draft_streaming( - chat_type=self.cfg.chat_type or None, - metadata=self.metadata, - chat_id=self.chat_id, + chat_id=self.chat_id, **probe_kwargs, ) except TypeError: - supported = self.adapter.supports_draft_streaming( - chat_type=self.cfg.chat_type or None, - metadata=self.metadata, - ) + supported = self.adapter.supports_draft_streaming(**probe_kwargs) except Exception: logger.debug("supports_draft_streaming probe raised", exc_info=True) supported = False - if not supported: - if transport == "draft": - logger.debug( - "Draft streaming requested but unsupported (chat=%s, type=%r) — " - "falling back to edit", - self.chat_id, self.cfg.chat_type, - ) - return False - return True + if not supported and transport == "draft": + logger.debug( + "Draft streaming requested but unsupported (chat=%s, type=%r) — " + "falling back to edit", + self.chat_id, self.cfg.chat_type, + ) + return bool(supported) def _resolve_native_streaming(self) -> bool: """Whether to use native streaming (adapter.send_stream_frame for ALL frames). @@ -184,31 +168,20 @@ class StreamTransportMixin: if probe is None: return False try: - supported = probe( - chat_type=self.cfg.chat_type or None, - metadata=self.metadata, - ) + return bool(probe(chat_type=self.cfg.chat_type or None, metadata=self.metadata)) except Exception: - logger.debug( - "supports_native_streaming probe raised", exc_info=True, - ) + logger.debug("supports_native_streaming probe raised", exc_info=True) return False - return bool(supported) async def _send_draft_frame(self, text: str) -> bool: """Emit one draft frame; any failure permanently disables drafts for this run. - Drafts have no message_id and clear on the client when the final - sendMessage lands. + Drafts have no message_id and clear on the client when the final sendMessage lands. """ if self._draft_id is None: # Should never happen (set in tandem with _use_draft_streaming in run()). self._use_draft_streaming = False return False - # Every frame must carry the same reply_to_message_id the final send - # gets from _metadata_for_send: the relay adapter keys draft/seal state - # on it, else the final can't find the open stream (flat DMs have no - # thread metadata and would key on the bare chat). try: result = await self.adapter.send_draft( chat_id=self.chat_id, @@ -217,34 +190,28 @@ class StreamTransportMixin: metadata=self._draft_metadata(), ) except Exception as e: - logger.debug( - "send_draft raised, disabling draft transport for this run: %s", e, - ) - self._draft_failures += 1 - self._use_draft_streaming = False - return False - if not getattr(result, "success", False): + logger.debug("send_draft raised, disabling draft transport for this run: %s", e) + else: + if getattr(result, "success", False): + self._last_sent_text = text # parity with the edit-based no-op skip + return True logger.debug( "send_draft returned success=False, disabling draft transport: %s", getattr(result, "error", "unknown"), ) - self._draft_failures += 1 - self._use_draft_streaming = False - return False - self._last_sent_text = text # parity with the edit-based no-op skip - return True + self._draft_failures += 1 + self._use_draft_streaming = False + return False async def _abandon_native_stream(self) -> None: """Seal an orphaned draft stream in place on turn death (stale exit / cancel). - Otherwise the message keeps its live indicator forever and the - adapter's armed interception state leaks into the next turn. Never - sets delivery flags — the gateway's normal paths own what happens next. + Else the live indicator stays forever and the adapter's armed interception + state leaks into the next turn. Never sets delivery flags. """ if not self._use_draft_streaming: return - abandon = getattr(type(self.adapter), "abandon_open_draft", None) - if abandon is None: + if getattr(type(self.adapter), "abandon_open_draft", None) is None: return try: await self.adapter.abandon_open_draft( @@ -255,17 +222,16 @@ class StreamTransportMixin: except Exception as e: logger.debug("abandon_open_draft failed (best-effort): %s", e) + def _has_real_preview(self) -> bool: + """A real (editable, deletable) preview message id is on screen.""" + return bool(self._message_id) and self._message_id != "__no_edit__" + def _should_send_fresh_final(self) -> bool: """True when fresh-final is enabled and a real preview has been visible ≥ threshold.""" threshold = getattr(self.cfg, "fresh_final_after_seconds", 0.0) or 0.0 - if threshold <= 0: + if threshold <= 0 or not self._has_real_preview() or self._message_created_ts is None: return False - if not self._message_id or self._message_id == "__no_edit__": - return False - if self._message_created_ts is None: - return False - age = time.monotonic() - self._message_created_ts - return age >= threshold + return time.monotonic() - self._message_created_ts >= threshold def _track_preview_id(self, message_id: Optional[str]) -> None: """Record a real preview message id for finalization cleanup.""" @@ -289,7 +255,7 @@ class StreamTransportMixin: False when there's no real preview, no hook, or on any error. """ - if not self._message_id or self._message_id == "__no_edit__": + if not self._has_real_preview(): return False fn = getattr(self.adapter, "prefers_fresh_final_streaming", None) if fn is None: @@ -297,8 +263,8 @@ class StreamTransportMixin: try: try: # chat_id lets relay adapters decide via THIS chat's platform; - # otherwise a Slack-primary relay misroutes fronted chats - # through the fresh-send lane (duplicates: no delete op). + # otherwise a Slack-primary relay misroutes fronted chats through the + # fresh-send lane (duplicates: no delete op). result = fn(text, metadata=self.metadata, chat_id=self.chat_id) except TypeError: try: @@ -314,13 +280,11 @@ class StreamTransportMixin: async def _try_fresh_final(self, text: str, *, is_turn_final: bool = True) -> bool: """Send ``text`` as a fresh message and best-effort delete the preview(s). - Returns False on any failure so the caller falls back to edit. - ``is_turn_final=False`` (interim segment at a tool boundary) leaves the - final-delivery flag unset so the gateway still delivers the real answer. + False on any failure so the caller falls back to edit. ``is_turn_final=False`` + (interim segment) leaves the final-delivery flag unset. """ - # Replacing every tracked preview is only sound while ``text`` holds - # the whole answer; after a split, deleting sealed heads would erase - # delivered text — take the edit path instead. + # Replacing every preview is only sound while ``text`` holds the whole answer; + # after a split, deleting sealed heads would erase delivered text. if self._turn_split_delivery: return False stale_ids = self._stale_preview_ids() @@ -339,13 +303,7 @@ class StreamTransportMixin: # Best-effort preview cleanup; never delete the message just sent. await self._delete_previews(stale_ids, skip=new_message_id, label="Fresh-final") self._preview_message_ids = set() - if new_message_id: - self._message_id = new_message_id - self._message_created_ts = time.monotonic() - else: - # No id returned: sentinel so we never try to edit it. - self._message_id = "__no_edit__" - self._message_created_ts = None + self._adopt_message_id(new_message_id) self._already_sent = True self._last_sent_text = text if is_turn_final: @@ -353,31 +311,35 @@ class StreamTransportMixin: self._record_turn_final_payload(text) return True + def _adopt_message_id(self, message_id) -> None: + """Retarget edits at ``message_id``; None → "__no_edit__" sentinel so we never edit it.""" + if message_id: + self._message_id = message_id + self._message_created_ts = time.monotonic() + else: + self._message_id = "__no_edit__" + self._message_created_ts = None + async def _send_or_edit( self, text: str, *, finalize: bool = False, is_turn_final: bool = True, ) -> bool: """Send or edit the streaming message; True if delivered. - ``finalize`` marks the last edit of a streaming sequence. Callers such - as the overflow split loop use the result to decide whether to advance. - Transport order: native frame → draft frame → edit existing → first - send; a transport returns None to fall through to the next. + ``finalize`` marks the last edit of a streaming sequence. Transport order: + native frame → draft frame → edit existing → first send; a transport returns + None to fall through to the next. """ text = self._clean_for_display(text) - # Stream-is-the-message draft frames must stay prefix-stable: a closing - # ``` appended to a mid-code-block frame makes frame N not a prefix of - # N+1 and the connector re-appends the whole snapshot. The final - # message is still fence-closed below. + # Stream-is-the-message draft frames must stay prefix-stable: a closing ``` + # on a mid-code-block frame makes frame N not a prefix of N+1 and the + # connector re-appends the whole snapshot. The final is still fence-closed. pre_fence_text = text text = ensure_closed_code_fences(text) # A bare cursor renders as a stray tofu box on some clients. - visible_stripped = text - if self.cfg.cursor: - visible_stripped = visible_stripped.replace(self.cfg.cursor, "") - visible_stripped = visible_stripped.strip() + visible_stripped = (text.replace(self.cfg.cursor, "") if self.cfg.cursor else text).strip() if not visible_stripped: - # Native streams MUST still get a finalize frame (placeholder) to - # close the thinking bubble, e.g. for a MEDIA-only response. + # Native streams MUST still get a finalize frame (placeholder) to close + # the thinking bubble, e.g. for a MEDIA-only response. if finalize and self._use_native_streaming and self._native_stream_opened: try: if await self._send_frame("✅", finalize=True): @@ -385,11 +347,8 @@ class StreamTransportMixin: except Exception as e: logger.debug("Finalize empty stream failed: %s", e) return True # cursor-only / whitespace-only update - if not text.strip(): - return True # nothing to send is "success" - # Don't open a new message for 1-2 tokens + cursor (rapid tool-calling): - # if the cursor-strip edit is then rate-limited, "X ▉" stays forever. - # Only first sends are gated. + # Don't open a new message for 1-2 tokens + cursor (rapid tool-calling): if + # the cursor-strip edit is then rate-limited, "X ▉" stays forever. if ( self._message_id is None and self.cfg.cursor @@ -421,11 +380,12 @@ class StreamTransportMixin: logger.error("Stream send/edit error: %s", e) return False - async def _native_push(self, text: str, *, finalize: bool, is_turn_final: bool) -> Optional[bool]: + async def _native_push( + self, text: str, *, finalize: bool, is_turn_final: bool, + ) -> Optional[bool]: """Native streaming: every frame goes through send_stream_frame(). - The adapter's send/edit paths are not touched in this mode. Lazy - re-seed here after a boundary closed the stream. Returns None when + Lazy re-seed here after a boundary closed the stream. Returns None when native was disabled (seed/frame failure) so the caller falls through. """ if not self._native_stream_opened and text: @@ -447,23 +407,18 @@ class StreamTransportMixin: if not self._use_native_streaming: return None - # WeCom renders each finalize as a separate bubble: only the - # turn-final and boundaries close the stream, not segment breaks. - if finalize and not is_turn_final: - finalize = False + # WeCom renders each finalize as a separate bubble: only the turn-final and + # boundaries close the stream, not segment breaks. + finalize = finalize and is_turn_final if not finalize and text == self._last_sent_text: return True # unchanged — skip - # Mark a finalize frame delivered OPTIMISTICALLY, before the ack wait: - # the bytes hit the wire (and WeCom renders them) before the ack, so a - # gateway join-cancel during the ack wait must not strand - # final_content_delivered=False and cause a duplicate normal send - # (docs/rca-wecom-stream-final-ack-timeout-duplicate.md). A definitive - # dispatch failure rolls the mark back below. Residual window (cancel - # between mark and wire write, sub-ms) is accepted. + # Mark a finalize frame delivered OPTIMISTICALLY, before the ack wait: WeCom + # renders the bytes before the ack, so a gateway join-cancel mid-wait must not + # strand final_content_delivered=False and duplicate the send (docs/rca-wecom- + # stream-final-ack-timeout-duplicate.md). A definitive failure rolls it back. if finalize: - # Recorded so a stale/partial frame can't suppress the corrective send. - self._mark_final_delivered(record=text) + self._mark_final_delivered(record=text) # recorded: stale frame can't suppress try: ok = await self._send_frame(text, finalize=finalize) except Exception as e: @@ -483,13 +438,12 @@ class StreamTransportMixin: self._final_response_sent = False self._final_content_delivered = False self._delivered_final_text = None - # Subsequent frames take the edit/send fallback; the adapter marks the - # chat expired so it doesn't retry the dead stream. + # Subsequent frames take the edit/send fallback; the adapter marks the chat + # expired so it doesn't retry the dead stream. self._use_native_streaming = False - # Best-effort close of an opened bubble (seed frame has zero length but - # still opens it — hence _native_stream_opened, not pushed_len). DO NOT - # mark delivered: the frame closes the bubble but WeCom may not render - # the content (errcode 6000 race). + # Best-effort close of an opened bubble (the seed frame has zero length but + # still opens it). DO NOT mark delivered: the frame closes the bubble but + # WeCom may not render the content (errcode 6000 race). if self._native_stream_opened: try: await self._send_frame(text, finalize=True) @@ -503,12 +457,9 @@ class StreamTransportMixin: ) -> Optional[bool]: """Draft frame while no message_id exists; None = not applicable / drafts just failed. - Drafts have no message_id: the final answer goes through the regular - send (which clears the draft client-side), so drafts are skipped when - finalizing. Exception: stream-is-the-message adapters keep ONE stream - per turn, so a segment-break finalize must NOT become a real send - (seal interception would seal the stream at every tool boundary); - only got_done seals. + Drafts are skipped when finalizing (the real send clears the draft), EXCEPT + stream-is-the-message adapters keep ONE stream per turn: a segment-break + finalize must not become a real send (it would seal at every tool boundary). """ stream_is_msg = self._stream_is_message() if finalize and not (stream_is_msg and not is_turn_final): @@ -521,11 +472,9 @@ class StreamTransportMixin: frame_text = frame_text[: -len(self.cfg.cursor)] if frame_text == self._last_sent_text: return True - if await self._send_draft_frame(frame_text): - # Deliberately NOT _already_sent: the gateway's fallback final send - # must still fire so the user gets a real message. - return True - return None + # Deliberately NOT _already_sent on success: the gateway's fallback final + # send must still fire so the user gets a real message. + return True if await self._send_draft_frame(frame_text) else None async def _first_send(self, text: str, *, finalize: bool) -> bool: """First send, threaded to the user's message (correct topic/thread).""" @@ -538,36 +487,29 @@ class StreamTransportMixin: if not result.success: self._edit_supported = False return False - if result.message_id: - self._message_id = result.message_id - self._message_created_ts = time.monotonic() - self._track_preview_ids_from_result(result) - else: - self._edit_supported = False self._already_sent = True self._last_sent_text = text - if not result.message_id: - self._fallback_prefix = self._visible_prefix() - self._fallback_final_send = True - # Sentinel: no editable id, don't re-enter first-send on every - # delta/tool boundary. + if result.message_id: + self._adopt_message_id(result.message_id) + self._track_preview_ids_from_result(result) + else: + # No editable id: fallback mode + sentinel so we don't re-enter first-send. + self._enter_fallback_mode(self._visible_prefix()) self._message_id = "__no_edit__" self._notify_new_message() return True async def _edit_existing(self, text: str, *, finalize: bool, is_turn_final: bool) -> bool: """Edit the live preview (or replace it via fresh-final when finalizing).""" - # REQUIRES_EDIT_FINALIZE adapters need the finalize=True edit even - # when unchanged; everyone else short-circuits. + # REQUIRES_EDIT_FINALIZE adapters need the finalize=True edit even when + # unchanged; everyone else short-circuits. if text == self._last_sent_text and not (finalize and self._adapter_requires_finalize): return True - # Fresh-final: replace a long-lived preview with a fresh message - # (timestamp reflects completion), or whenever the adapter prefers it - # (Telegram's send path renders richer markdown than its edit path). - # An explicit hook returning False must NOT be overridden by the time - # threshold — on Telegram both messages would stay on screen since the - # delete is best-effort. Check the CLASS (MagicMock auto-creates - # attrs) plus instance __dict__ (test doubles assign the hook explicitly). + # Fresh-final: replace a long-lived preview with a fresh message, or whenever + # the adapter prefers it (Telegram's send path renders richer markdown). An + # explicit hook returning False must NOT be overridden by the time threshold + # (delete is best-effort; both messages would stay on screen). Check the + # CLASS (MagicMock auto-creates attrs) plus instance __dict__ (test doubles). has_prefers_hook = ( hasattr(type(self.adapter), "prefers_fresh_final_streaming") or "prefers_fresh_final_streaming" in getattr(self.adapter, "__dict__", {}) @@ -583,18 +525,19 @@ class StreamTransportMixin: message_id=self._message_id, content=text, finalize=finalize, ) if not result.success: - return await self._on_edit_failure(result, text, finalize=finalize, is_turn_final=is_turn_final) + return await self._on_edit_failure( + result, text, finalize=finalize, is_turn_final=is_turn_final, + ) self._already_sent = True self._track_preview_ids_from_result(result) - # Oversized edit split across continuations: message_id is now the - # LAST continuation, which holds only the final chunk — retarget edits - # and reset skip-if-same. getattr keeps SimpleNamespace test mocks working. + # Oversized edit split across continuations: message_id is now the LAST + # continuation, which holds only the final chunk — retarget edits and reset + # skip-if-same. getattr keeps SimpleNamespace test mocks working. continuation_ids = getattr(result, "continuation_message_ids", ()) or () if continuation_ids and result.message_id and result.message_id != self._message_id: self._last_edit_overflowed = True self._turn_split_delivery = True - self._message_id = str(result.message_id) - self._message_created_ts = time.monotonic() + self._adopt_message_id(str(result.message_id)) self._last_sent_text = "" self._notify_new_message() else: @@ -602,49 +545,54 @@ class StreamTransportMixin: self._flood_strikes = 0 return True - async def _on_edit_failure(self, result, text: str, *, finalize: bool, is_turn_final: bool) -> bool: + def _enter_fallback_mode(self, prefix: str) -> None: + """Edits are over for this stream: send only the missing tail at got_done.""" + self._fallback_prefix = prefix + self._fallback_final_send = True + self._edit_supported = False + self._already_sent = True + + async def _on_edit_failure( + self, result, text: str, *, finalize: bool, is_turn_final: bool, + ) -> bool: """Classify a failed edit: partial overflow, flood backoff, or fallback mode. - Returns the _send_or_edit result (always False here; the caller's - finalize path may still deliver the tail via _send_fallback_final). + Always False; the caller's finalize path may still deliver the tail. """ - immediate_final_fallback = False + turn_final = finalize and is_turn_final if ( - finalize - and is_turn_final + turn_final and self.cfg.cursor and self._last_sent_text.endswith(self.cfg.cursor) and self._visible_prefix() == text ): - # Cosmetic final edit was rate-limited but the full answer is - # already on screen (cursor stuck): mark delivered so the gateway - # doesn't send it twice, and record the on-screen payload. + # Cosmetic final edit was rate-limited but the full answer is already on + # screen (cursor stuck): mark delivered so the gateway doesn't send it + # twice, and record the on-screen payload. self._final_content_delivered = True self._record_turn_final_payload(text) raw_response = getattr(result, "raw_response", None) if isinstance(raw_response, dict) and raw_response.get("partial_overflow"): - # Some overflow chunks landed but not the whole response: preserve - # the visible prefix so got_done sends the missing tail. + # Some overflow chunks landed but not the whole response: preserve the + # visible prefix so got_done sends the missing tail. self._message_id = str( raw_response.get("last_message_id") or result.message_id or self._message_id ) delivered_prefix = raw_response.get("delivered_prefix") if isinstance(delivered_prefix, str) and delivered_prefix: self._last_sent_text = delivered_prefix - self._fallback_prefix = delivered_prefix self._fallback_preserve_partial_messages = text.startswith(delivered_prefix) + self._enter_fallback_mode(delivered_prefix) else: - self._fallback_prefix = self._visible_prefix() self._fallback_preserve_partial_messages = False - self._fallback_final_send = True - self._edit_supported = False - self._already_sent = True + self._enter_fallback_mode(self._visible_prefix()) if getattr(result, "continuation_message_ids", ()): self._notify_new_message() return False - # Flood control: adaptive backoff (double the interval); disable edits - # only after _MAX_FLOOD_STRIKES in a row. + # Flood control: adaptive backoff (double the interval); disable edits only + # after _MAX_FLOOD_STRIKES in a row. + immediate_final_fallback = False if self._is_flood_error(result): self._flood_strikes += 1 self._current_edit_interval = min(self._current_edit_interval * 2, 10.0) @@ -653,9 +601,7 @@ class StreamTransportMixin: self._flood_strikes, self._MAX_FLOOD_STRIKES, self._current_edit_interval, ) immediate_final_fallback = ( - finalize - and is_turn_final - and getattr(self.adapter, "FALLBACK_ON_FINAL_EDIT_FLOOD", False) is True + turn_final and getattr(self.adapter, "FALLBACK_ON_FINAL_EDIT_FLOOD", False) is True ) if self._flood_strikes < self._MAX_FLOOD_STRIKES and not immediate_final_fallback: self._last_edit_time = time.monotonic() # honor the new interval @@ -663,14 +609,10 @@ class StreamTransportMixin: if immediate_final_fallback: logger.debug("Turn-final edit hit flood control; entering fallback immediately") - # Fallback mode: send only the missing tail at got_done. logger.debug("Edit failed (strikes=%d), entering fallback mode", self._flood_strikes) - self._fallback_prefix = self._visible_prefix() - self._fallback_final_send = True - self._edit_supported = False - self._already_sent = True - # A turn-final flood skips the cosmetic cursor strip: it would burn the - # same flood budget and delay the answer. + self._enter_fallback_mode(self._visible_prefix()) + # A turn-final flood skips the cosmetic cursor strip: it would burn the same + # flood budget and delay the answer. if not immediate_final_fallback: await self._try_strip_cursor() return False diff --git a/gateway/stream_dispatch.py b/gateway/stream_dispatch.py index 0a6b63a53f..f2f7d4f3d9 100644 --- a/gateway/stream_dispatch.py +++ b/gateway/stream_dispatch.py @@ -1,14 +1,11 @@ """Adapter-driven dispatch of structured stream events to a delivery sink. -``GatewayEventDispatcher`` holds an adapter, the stream consumer (sink) and the -resolved per-channel presentation settings, and routes each typed event -(gateway/stream_events.py) through the adapter's render hooks. Message events -flow into the consumer; tool events are formatted by the adapter — which may -return None to *eat* them on platforms without tool chrome — and enqueued onto -the same tool-progress queue the gateway drains, so the two paths never race. - -No platform knowledge and no asyncio: a thin synchronous router callable from -the agent's worker thread, exactly like the callbacks it replaced. +``GatewayEventDispatcher`` routes each typed event (gateway/stream_events.py) +through the adapter's render hooks: message events flow into the consumer; tool +events are formatted by the adapter — which may return None to *eat* them on +platforms without tool chrome — and enqueued onto the same tool-progress queue the +gateway drains, so the two paths never race. Synchronous, no asyncio: callable +from the agent's worker thread. """ from __future__ import annotations @@ -26,14 +23,13 @@ logger = logging.getLogger("gateway.stream_events") class GatewayEventDispatcher: """Route typed stream events through an adapter onto a delivery sink. - adapter: provides ``render_message_event`` / ``format_tool_event`` - (BasePlatformAdapter defaults reproduce legacy behavior). + adapter: provides ``render_message_event`` / ``format_tool_event``. sink: the GatewayStreamConsumer; None when streaming is disabled (message events are dropped — the final response still goes out normally). enqueue_tool_line: puts a rendered tool-progress line on the gateway's progress queue; None when tool progress is disabled for the channel. - tool_mode: "all" / "new" / "verbose" / "off". - preview_max_len: resolved ``tool_preview_length`` (0 = no cap in verbose). + tool_mode: "all" / "new" / "verbose" / "off". preview_max_len: resolved + ``tool_preview_length`` (0 = no cap in verbose). on_long_tool / on_notice: optional hooks so the gateway owns the "should I surface this here?" decision. """ @@ -66,6 +62,8 @@ class GatewayEventDispatcher: logger.debug("stream-event dispatch error", exc_info=True) def _dispatch(self, event: StreamEvent) -> None: + # ToolCallFinished: no chrome on completion (only "started" is rendered); + # completion only drives onboarding hints (LongToolHint). if isinstance(event, (MessageChunk, MessageStop, Commentary)): if self.sink is not None: self.adapter.render_message_event(event, self.sink) @@ -75,8 +73,6 @@ class GatewayEventDispatcher: self._on_long_tool(event) elif isinstance(event, GatewayNotice) and self._on_notice is not None: self._on_notice(event) - # ToolCallFinished: no chrome on completion (only "started" is rendered); - # completion only drives onboarding hints (LongToolHint). def _dispatch_tool_call(self, event: ToolCallChunk) -> None: if self.tool_mode == "off" or self._enqueue_tool_line is None: diff --git a/gateway/stream_events.py b/gateway/stream_events.py index ac0df25f28..86156a2cc3 100644 --- a/gateway/stream_events.py +++ b/gateway/stream_events.py @@ -1,23 +1,11 @@ """Structured streaming events — the agent→gateway delivery contract. A small typed vocabulary naming *what happened* without prescribing *how it is -delivered*. The agent emits these from its worker thread; the gateway's stream -consumer (``GatewayStreamConsumer``) is the single sink and the platform adapter -decides rendering (Telegram may stream a native draft; iMessage may drop tool -chrome). This replaced a fan of loosely-typed callbacks whose gateway side -decided both rendering and sending — the cause of tool-progress bubbles racing -the streaming draft. - -Plain frozen dataclasses: no behavior, no platform knowledge, no I/O, safe to -hand across the thread/async boundary. - -Invariants: - * Events describe *transport*, never *context*. Nothing here is persisted; - whatever the gateway chooses to "eat" must never diverge from the bytes in - the agent's message history, which the agent alone owns. - * Backward compatible by construction: the gateway adapts existing callbacks - into events at the boundary; adapters that don't opt in get identical - behavior via the base-class default. +delivered*: the agent emits these from its worker thread, ``GatewayStreamConsumer`` +is the single sink and the platform adapter decides rendering. Plain frozen +dataclasses: no behavior, no I/O, safe across the thread/async boundary. Events +describe *transport*, never *context* — whatever the gateway "eats" must never +diverge from the agent-owned message history. """ from __future__ import annotations @@ -28,10 +16,7 @@ from typing import Any, Dict, Optional, Union @dataclass(frozen=True) class MessageChunk: - """A delta of streamed assistant text (reasoning/think-block content is - filtered upstream and never arrives here); the consumer accumulates chunks - and renders progressively (native draft on Telegram DMs, edit-in-place - elsewhere).""" + """A delta of streamed assistant text (think-block content is filtered upstream).""" text: str @@ -39,9 +24,9 @@ class MessageChunk: class MessageStop: """The current assistant text segment is complete. - ``final`` is True only for the terminal stop of the turn; an intermediate - stop (text → tool call → more text) carries ``final=False`` so the consumer - finalizes the current bubble and starts a fresh segment below tool chrome. + ``final`` is True only for the terminal stop of the turn; an intermediate stop + (text → tool call → more text) makes the consumer start a fresh segment below + tool chrome. """ final: bool = False @@ -58,8 +43,7 @@ class ToolCallChunk: tool_name: str preview: Optional[str] = None args: Optional[Dict[str, Any]] = None - # Monotonic per-turn index: correlates a finish with its start and lets - # "new"-mode dedup work without the consumer tracking call order. + # Monotonic per-turn index: correlates a finish with its start. index: int = 0 @@ -67,8 +51,7 @@ class ToolCallChunk: class ToolCallFinished: """A tool invocation completed. Tool *output* never travels here (it is history). - The gateway uses it to clear/settle a progress bubble and to drive one-time - onboarding hints (e.g. suggest /verbose after a long tool run). + Drives progress-bubble settling and one-time onboarding hints (LongToolHint). """ tool_name: str duration: float = 0.0 # wall-clock seconds @@ -80,8 +63,7 @@ class ToolCallFinished: class LongToolHint: """One-shot onboarding nudge when a tool runs longer than the threshold. - The gateway (not the agent) gates it on platform capability (the /verbose - command must be usable) and on the user not having seen it before. + The gateway gates it on platform capability (/verbose usable) and first-time use. """ tool_name: str = "" duration: float = 0.0 @@ -91,8 +73,8 @@ class LongToolHint: class GatewayNotice: """A gateway-originated control message. - ``kind`` is a stable string adapters switch on (``"restart"`` / ``"online"`` - / ``"long_run"`` / …); ``text`` is the default rendering. + ``kind`` is a stable string adapters switch on (``"restart"`` / ``"online"`` / + ``"long_run"`` / …); ``text`` is the default rendering. """ kind: str text: str = "" diff --git a/gateway/streaming_tts_consumer.py b/gateway/streaming_tts_consumer.py index e25982e713..865d4dcb7e 100644 --- a/gateway/streaming_tts_consumer.py +++ b/gateway/streaming_tts_consumer.py @@ -1,39 +1,24 @@ """Gateway streaming-TTS consumer — LLM deltas to adapter PCM audio sink. -Bridges the synchronous agent ``stream_delta_callback`` (worker thread) to a -voice-capable adapter's streaming-audio contract so playback begins while the -LLM is still generating. - -Lifecycle:: - - consumer = StreamingTTSConsumer(adapter, chat_id, tts_config, loop, metadata) - agent.stream_delta_callback = consumer.on_delta # sync, non-blocking - ... agent runs in executor ... - consumer.finish() # signal end-of-text - success = await consumer.wait_complete(timeout=10) - if consumer.suppress_whole_file: ... # skip whole-file auto-TTS - consumer.abort("cancelled") # idempotent cancellation - -``on_delta`` never blocks: it feeds a ``SentenceChunker`` and queues clauses on a -thread-safe ``queue.Queue``; the ``_run`` task on the gateway loop drains it, -synthesises via a ``StreamingTTSProvider`` and writes PCM to the adapter. State -is per instance (concurrent chats cannot cross-contaminate); abort is idempotent -and late chunks are dropped. Outcome contract: full success -> ``completed``; -failure before any audible output -> ``suppress_whole_file=False`` (gateway falls -back to whole-file TTS); failure after partial audio -> ``partial`` and -``suppress_whole_file=True`` (never replay the response from the beginning). +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). """ from __future__ import annotations import asyncio +import contextlib import logging import queue import threading from typing import Any, Dict, Optional from gateway.platforms.base import AudioFormat, StreamingTTSHandle -import contextlib logger = logging.getLogger("gateway.streaming_tts_consumer") @@ -60,94 +45,71 @@ class StreamingTTSConsumer: self._chat_id = chat_id self._loop = loop self._metadata = 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( - sample_rate=int(getattr(self._streamer, "sample_rate", AudioFormat.sample_rate)), - channels=int(getattr(self._streamer, "channels", AudioFormat.channels)), - sample_width=int(getattr(self._streamer, "sample_width", AudioFormat.sample_width)), - ) + 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() - # 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._completed = False - self._partial = False - self._aborted = False - self._finished = False - self._dropped = False - self._suppress_whole_file = False + 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 - @property - def active(self) -> bool: # usable streaming provider resolved - return self._streamer is not None + # 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()) - @property - def completed(self) -> bool: # streaming audio fully delivered - return self._completed - - @property - def partial(self) -> bool: # some audio was audible before a failure/drop - return self._partial - - @property - def audible(self) -> bool: # first PCM chunk has been written - return bool(self._handle and self._handle.audible) - - @property - def dropped(self) -> bool: # queue saturation dropped at least one clause - return self._dropped - - @property - def suppress_whole_file(self) -> bool: # gateway should skip whole-file TTS fallback - return self._suppress_whole_file - - @property - def done(self) -> bool: # async drain task has terminated - return self._task is not None and self._task.done() + def _enqueue_clauses(self, clauses, full_msg: str, *, log_errors: bool) -> None: + try: + for clause in clauses: + self._queue.put_nowait(clause) + except queue.Full: + self._dropped = True + logger.debug(full_msg) + except Exception: + if log_errors: + logger.debug("streaming TTS on_delta error", exc_info=True) def on_delta(self, text: str) -> None: """Receive a text delta from the agent. Non-blocking.""" if self._aborted or not self.active or self._finished: return - try: - for clause in self._chunker.feed(text): - self._queue.put_nowait(clause) - except queue.Full: - self._dropped = True - logger.debug("streaming TTS queue full, dropping clause") - except Exception: - logger.debug("streaming TTS on_delta error", exc_info=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``. - - The sentinel follows all flushed clauses so the drain loop has a - deterministic termination that cannot race a late ``on_delta``. + """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``. """ if self._finished: return self._finished = True if self._aborted or not self.active: return - try: - for clause in self._chunker.flush(): - self._queue.put_nowait(clause) - except queue.Full: - self._dropped = True - logger.debug("streaming TTS queue full while flushing tail") - except Exception: - pass + 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 a clause if full. while True: try: @@ -167,12 +129,8 @@ class StreamingTTSConsumer: return self._task def _settle(self, *, failed: bool) -> None: - """Set the outcome flags from what was audible. - - Never report completion after a failure or a dropped clause; keep - suppression whenever audio was audible so the gateway does not replay - the response from the beginning. - """ + """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 self._completed = audible and not degraded @@ -189,7 +147,7 @@ class StreamingTTSConsumer: return try: self._handle = await self._adapter.begin_streaming_tts( - self._chat_id, self._audio_format, metadata=self._metadata, + self._chat_id, self._audio_format, metadata=self._metadata ) except Exception as exc: logger.debug("begin_streaming_tts failed: %s", exc) @@ -197,7 +155,6 @@ class StreamingTTSConsumer: return if self._handle is None: return - self._suppress_whole_file = False try: while not self._aborted: @@ -216,7 +173,6 @@ class StreamingTTSConsumer: 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) @@ -230,11 +186,9 @@ class StreamingTTSConsumer: logger.warning("streaming TTS consumer error: %s", exc) await self._safe_abort(str(exc)) finally: - try: + with contextlib.suppress(Exception): while not self._queue.empty(): self._queue.get_nowait() - except Exception: - pass async def _synthesise_and_write(self, clause: str) -> None: """Synthesise one clause via the streamer and write PCM chunks.""" @@ -247,9 +201,7 @@ class StreamingTTSConsumer: while True: # next() runs in a thread so a blocking provider never stalls the loop. chunk = await asyncio.to_thread(next, iterator, _DONE) - if chunk is _DONE: - return - if self._aborted or self._handle.aborted: + if chunk is _DONE or self._aborted or self._handle.aborted: return if not chunk: continue @@ -274,9 +226,8 @@ class StreamingTTSConsumer: if self._handle is None: return try: - await self._adapter.abort_streaming_tts(self._handle, error=reason) - except Exception: - pass + with contextlib.suppress(Exception): + await self._adapter.abort_streaming_tts(self._handle, error=reason) finally: if self._handle: self._handle.aborted = True @@ -287,8 +238,7 @@ class StreamingTTSConsumer: if self._aborted: return self._aborted = True - # The _ABORT sentinel is load-bearing and must reach the queue even when - # the bounded queue is full: evict an item to make room. + # The load-bearing _ABORT sentinel must reach the queue even when full: evict to make room. for _attempt in range(3): try: self._queue.put_nowait(_ABORT) diff --git a/gateway/turn_context.py b/gateway/turn_context.py index 192a7766ae..477c8e8b04 100644 --- a/gateway/turn_context.py +++ b/gateway/turn_context.py @@ -1,20 +1,10 @@ -"""Per-turn context shared between ``GatewayRunner._run_agent_inner`` and -``TurnRunner`` (gateway/run.py). +"""Per-turn context shared between ``GatewayRunner._run_agent_inner`` and ``TurnRunner``. -``_run_agent_inner`` once defined its tool-progress plumbing and ``run_sync`` as -nested closures over ~20 locals. ``TurnContext`` is the extraction seam: each -closed-over local is a field here, so the closure bodies moved onto ``TurnRunner`` -methods unchanged modulo ``name`` -> ``ctx.name`` rewrites. - -Invariants: -- Fields are written once by ``_run_agent_inner`` while wiring the turn (a few are - assigned onto the ctx slightly after construction, at the original binding sites). -- The original closures never rebound captured names except ``message`` (formerly - ``nonlocal``): rebind sites now write ``ctx.message`` and the outer body reads it. - Other mutable state keeps the single-element-list containers so mutation stays - visible to the outer body through the shared objects. -- ``_run_still_current`` stays a callable (captures ``self``/``session_key``/ - ``run_generation``) so the extracted bodies remain byte-identical. +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``. """ from __future__ import annotations @@ -27,7 +17,7 @@ from typing import Any, Callable, List, Optional class TurnContext: """Closed-over locals of ``_run_agent_inner`` needed by ``TurnRunner``.""" - # --- read-only turn identity / wiring --- + # read-only turn identity / wiring source: Any = None _run_still_current: Callable[[], bool] = None # type: ignore[assignment] _live_status_adapter: Any = None @@ -36,32 +26,25 @@ class TurnContext: progress_mode: str = "off" progress_grouping: str = "grouped" tool_progress_enabled: bool = False - - # --- queues --- progress_queue: Any = None log_queue: Any = None - - # --- mutable single-element containers (shared with the outer body) --- + # mutable single-element containers (shared with the outer body) last_progress_msg: list = field(default_factory=lambda: [None]) last_tool: list = field(default_factory=lambda: [None]) last_was_terminal_block: list = field(default_factory=lambda: [False]) 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 --- + # 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 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) --- + # run_sync seam: the ex-``nonlocal`` turn message (rebindable) message: Optional[str] = None - - # --- turn parameters / config snapshots (read-only in run_sync) --- + # turn parameters / config snapshots (read-only in run_sync) history: Any = None context_prompt: Optional[str] = None channel_prompt: Optional[str] = None @@ -72,17 +55,14 @@ 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 - # event_message_id (the reply/thread anchor, which may be the replied-to message - # on Slack/Mattermost/Buzz or None for Telegram topics). Stamped as - # platform_message_id on the persisted user turn. + # 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. 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 when the turn was self-injected - # (MessageEvent.internal), e.g. "internal_notification". DB-only presentation - # metadata; never sent to the provider. + # 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. persist_user_display_kind: Optional[str] = None user_config: Any = None enabled_toolsets: Any = None @@ -90,39 +70,32 @@ 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 --- + # lazy-imported callables captured from the outer body AIAgent: Any = None resolve_display_setting: Any = None - - # --- mutable holder cells (shared-list pattern) --- + # 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]) streaming_tts_consumer_holder: list = field(default_factory=lambda: [None]) - - # --- voice-ack wiring --- + # voice-ack wiring _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 (published at original binding sites) _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) --- + # extracted sibling callbacks (bound TurnRunner methods 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 the adapter's - # native_task_cards_enabled()). ID-bearing lifecycle callbacks are published - # by TurnRunner so tool starts/completions correlate by tool-call ID. --- + # 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 _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 60e3441e39..91b7f2ee32 100644 --- a/gateway/turn_lease.py +++ b/gateway/turn_lease.py @@ -1,36 +1,15 @@ -"""Per-session turn lease — serializes the [load history → run → flush] region. +"""Per-session turn lease — serializes the [load history -> run -> flush] region. -Why: the gateway's busy guards are keyed by ROUTING KEY (adapter -``_active_sessions``, runner ``_running_agents``), but the durable transcript is -owned by SESSION_ID, and ``switch_session()`` makes key→id many-to-one (/resume -of a named session from a second chat, CLI-continuity rebinding, async-delegation -completion pinning, Telegram topic-binding tip-walks). Two routing keys on one -session_id ran concurrent turns on two agent objects that no per-key guard could -see; their flushes interleaved on one transcript (rows in completion order, -identity-marker dedup swallowing rows, a permanent ``user;user`` alternation -wedge that ``repair_message_sequence`` re-repaired forever). - -The lease serializes per RESOLVED session_id: acquired after resolution is final -(post switch_session/tip-walk), immediately before the transcript load, and -released in the dispatch layer's ``finally`` on every exit path. Same-key -messages never reach acquisition while a turn runs (routing-key guards hold -them), so the lock is uncontended except on the alias-key route, where the -second turn waits for the first's flush and logs one WARNING naming the session -and both keys (pairing with ``agent_runtime_helpers.note_turn_start``). - -Safety properties: -- Generation-scoped, identity-checked, idempotent release: a token records its - owner (routing key, run generation) and only the exact current holder frees - the lease — a stale unwind can never release a newer turn's lease. -- Fail-closed on timeout: a timed-out waiter raises :class:`TurnLeaseTimeoutError` - and must be rejected with a visible resend notice; it never runs against the - still-held lease. -- Bounded registry: eviction only removes idle (unheld, uncontended) entries. - -Known limits: a CLI process sharing the session via CLI-continuity is outside -any in-process lock (needs a DB-level lease); mid-turn compression rotation -leaves a small alias window, closed by :meth:`SessionTurnLeaseRegistry.rebind` -at the mid-turn binding-sync sites. +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`. """ import asyncio @@ -40,16 +19,12 @@ from typing import Dict, Optional logger = logging.getLogger(__name__) -# Cap on tracked per-session leases. Idle entries are evicted oldest-first; live -# leases never are, so a burst of distinct sessions may transiently exceed the cap -# rather than break serialization. +# Cap on tracked leases. Idle entries evict oldest-first; live leases never do, so a burst of +# distinct sessions may transiently exceed the cap rather than break serialization. DEFAULT_MAX_LEASES = 512 - -# Fallback wait (seconds) when the caller passes no positive timeout. The gateway -# carries this through its internal HERMES_TURN_LEASE_TIMEOUT bridge because lease -# contention is not agent inactivity. Fail-closed but short: never pin a -# sequential platform updater for minutes — a waiter that cannot acquire promptly -# is rejected with a resend notice, never authorized to run unserialized. +# Fallback wait (seconds) when the caller passes no positive timeout (bridged via +# HERMES_TURN_LEASE_TIMEOUT — lease contention is not agent inactivity). Fail-closed but short: +# never pin a sequential platform updater for minutes. DEFAULT_LEASE_WAIT = 5.0 @@ -62,17 +37,10 @@ class TurnLeaseTimeoutError(TimeoutError): caller must not enter the transcript load/run/flush region).""" def __init__( - self, - session_id: str, - *, - owner_key: str, - generation: int, - wait_seconds: float, + self, session_id: str, *, owner_key: str, generation: int, wait_seconds: float ) -> None: - self.session_id = session_id - self.owner_key = owner_key - self.generation = generation - self.wait_seconds = wait_seconds + 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})" @@ -80,25 +48,19 @@ class TurnLeaseTimeoutError(TimeoutError): class TurnLeaseToken: - """Handle returned by :meth:`SessionTurnLeaseRegistry.acquire`. - - Every token handed out is a held lease (timeouts raise instead); ``released`` - makes release idempotent. - """ + """Held-lease handle from :meth:`SessionTurnLeaseRegistry.acquire`; ``released`` makes + release idempotent.""" __slots__ = ("session_id", "owner_key", "generation", "released") def __init__(self, session_id: str, owner_key: str, generation: int) -> None: - self.session_id = session_id - self.owner_key = owner_key - self.generation = generation + self.session_id, self.owner_key, self.generation = session_id, owner_key, generation self.released = False def __repr__(self) -> str: # pragma: no cover - debug aid return ( - f"TurnLeaseToken(session_id={self.session_id!r}, " - f"owner_key={self.owner_key!r}, generation={self.generation}, " - f"released={self.released})" + f"TurnLeaseToken(session_id={self.session_id!r}, owner_key={self.owner_key!r}, " + f"generation={self.generation}, released={self.released})" ) @@ -119,11 +81,8 @@ class _SessionLease: class SessionTurnLeaseRegistry: - """Asyncio lease per resolved session_id serializing transcript turns. - - Process-local and single-event-loop by design — the same visibility scope as - the routing-key guards it extends. Call only from the gateway's event loop. - """ + """Asyncio lease per resolved session_id. Process-local, single-event-loop by design (same + visibility scope as the routing-key guards it extends); call only from the gateway loop.""" def __init__(self, max_entries: int = DEFAULT_MAX_LEASES) -> None: self._leases: Dict[str, _SessionLease] = {} @@ -141,8 +100,8 @@ class SessionTurnLeaseRegistry: return lease def _evict_idle(self) -> None: - """Drop oldest idle entries so a new lease fits under the cap. Never - evicts a held or contended lease — correctness beats the cap.""" + """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: return @@ -154,87 +113,55 @@ class SessionTurnLeaseRegistry: self._leases.pop(sid, None) async def acquire( - self, - session_id: str, - *, - owner_key: str, - generation: int, - timeout: Optional[float] = None, + self, session_id: str, *, owner_key: str, generation: int, timeout: Optional[float] = None ) -> Optional[TurnLeaseToken]: - """Acquire the turn lease for ``session_id``, waiting if held. - - Raises :class:`TurnLeaseTimeoutError` when the wait budget expires (the - caller must reject the turn). Returns None for a falsy ``session_id``. - """ + """Acquire the lease for ``session_id``, waiting if held. Raises + :class:`TurnLeaseTimeoutError` when the wait budget expires; None for a falsy id.""" if not session_id: return None wait = float(timeout) if timeout and timeout > 0 else DEFAULT_LEASE_WAIT token = TurnLeaseToken(session_id, owner_key, int(generation)) lease = self._get_or_create(session_id) - if lease.lock.locked(): logger.warning( - "turn lease contention on session %s: routing key %s (gen %s) " - "waiting behind in-flight turn held by routing key %s (gen %s, " - "held %.0fs) — two routing keys 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), + "turn lease contention on session %s: routing key %s (gen %s) waiting behind " + "in-flight turn held by routing key %s (gen %s, held %.0fs) — two routing keys " + "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, ) - - # Lock.release() wakes a waiter while leaving the lock momentarily - # unlocked. Count every in-progress acquire across that handoff so - # eviction cannot orphan the old lock and create a second lock for the - # same session — even apparently-uncontended ones, since wait_for() may - # schedule them before the underlying lock coroutine runs. + # 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 + # lock and create a second lock for the same session. lease.pending_acquires += 1 try: await asyncio.wait_for(lease.lock.acquire(), timeout=wait) except asyncio.TimeoutError: logger.error( - "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 " + "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 finally: lease.pending_acquires -= 1 - - # Lock held and no await before holder publication, so the lease cannot - # become evictable after the pending count is cleared. + # Lock held and no await before holder publication, so the lease cannot become + # evictable after the pending count is cleared. lease.holder = token lease.acquired_at = lease.last_used = time.time() return token def rebind(self, token: Optional[TurnLeaseToken], new_session_id: str) -> bool: - """Alias a HELD lease onto ``new_session_id`` after mid-turn rotation. - - Compression can rotate the durable session_id mid-turn (session-hygiene - pre-compression, in-agent compression); the flush then targets the NEW - id, so the serialization boundary must follow it or an alias key - resolving the new id could start a turn the lease never sees. - - Mechanism: the SAME ``_SessionLease`` is registered under the new id (the - old mapping stays until idle and evicted), so acquirers on either id - serialize on one lock — no lock state moves. Only the current holder can - rebind (identity-checked like release), and the token follows so release - frees the shared object. - - Edge: if the new id already has a live lease of its own, the two domains - cannot be merged mid-wait — log loudly and keep the token on the old id. - Fail-open, never deadlock: a holder cannot wait mid-turn. + """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). """ if ( token is None @@ -246,35 +173,25 @@ class SessionTurnLeaseRegistry: lease = self._leases.get(token.session_id) if lease 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: logger.warning( - "turn lease rebind blocked: session %s rotated to %s mid-turn " - "(holder: routing key %s gen %s) but the target session's " - "lease is already live (holder: routing key %s 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, + "turn lease rebind blocked: session %s rotated to %s mid-turn (holder: routing key " + "%s gen %s) but the target session's lease is already live (holder: routing key %s " + "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, ) return False - self._leases[new_session_id] = lease lease.last_used = time.time() token.session_id = new_session_id return True def release(self, token: Optional[TurnLeaseToken]) -> bool: - """Release ``token``'s lease. Idempotent; ownership-checked. - - True only when this exact token was the current holder. A re-release or - a stale token whose slot went to a newer turn are safe no-ops. - """ + """Release ``token``'s lease. Idempotent; True only when this exact token was the current + holder (a re-release or a stale token whose slot went to a newer turn is a safe no-op).""" if token is None or token.released: return False token.released = True @@ -283,16 +200,12 @@ class SessionTurnLeaseRegistry: 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, + "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 = None - lease.acquired_at = 0.0 - lease.last_used = time.time() + lease.holder, lease.acquired_at, lease.last_used = None, 0.0, time.time() if lease.lock.locked(): lease.lock.release() return True diff --git a/gateway/whatsapp_identity.py b/gateway/whatsapp_identity.py index 5271f45f4b..eda5975bbd 100644 --- a/gateway/whatsapp_identity.py +++ b/gateway/whatsapp_identity.py @@ -1,11 +1,9 @@ """Shared helpers for canonicalising WhatsApp sender identity. -The bridge can surface the same human as a LID (``999...@lid``) or a phone -JID (``1555...@s.whatsapp.net``) within one conversation. Authorisation -(:mod:`gateway.run`) and session keys (:mod:`gateway.session`) both resolve -those aliases here so the two paths can never drift apart. Plugins needing -per-sender behaviour on WhatsApp should use :func:`canonical_whatsapp_identifier` -so their bookkeeping lines up with Hermes' own session keys. +The bridge can surface one human as a LID (``999...@lid``) or a phone JID +(``1555...@s.whatsapp.net``) within one conversation. Authorisation (:mod:`gateway.run`) and +session keys (:mod:`gateway.session`) both resolve aliases here so they never drift apart; +plugins should use :func:`canonical_whatsapp_identifier` to line up with Hermes' session keys. """ from __future__ import annotations @@ -29,72 +27,51 @@ _BARE_PHONE_RE = re.compile(r"^\+?[\d\s().\-]+$") def normalize_whatsapp_identifier(value: str) -> str: - """Strip JID/LID/device/plus syntax down to the bare numeric identifier. - - ``"60123456789:47@s.whatsapp.net"``, ``"60123456789@lid"`` and - ``"+60123456789"`` all become ``"60123456789"``. + """Strip JID/LID/device/plus syntax down to the bare numeric identifier: + ``"6012:47@s.whatsapp.net"``, ``"6012@lid"`` and ``"+6012"`` all become ``"6012"``. """ - return ( - str(value or "") - .strip() - .replace("+", "", 1) - .split(":", 1)[0] - .split("@", 1)[0] - ) + return str(value or "").strip().replace("+", "", 1).split(":", 1)[0].split("@", 1)[0] def to_whatsapp_jid(value: str) -> str: """Normalize an *outbound* target to a bridge-safe JID (inverse of normalize). Baileys' ``jidDecode`` crashes on a bare phone number, so bare phones become - ``@s.whatsapp.net``; ``user:device@domain`` collapses to - ``user@domain``; anything else already carrying ``@`` or not recognizable - as a phone is returned unchanged so the bridge can surface a real error. - Returns ``""`` for empty input. + ``@s.whatsapp.net``; ``user:device@domain`` collapses to ``user@domain``; anything + else is returned unchanged so the bridge can surface a real error. ``""`` for empty input. """ if not value: return "" - normalized = str(value).strip() if ":" in normalized and "@" in normalized: prefix, _, domain = normalized.partition("@") normalized = f"{prefix.split(':', 1)[0]}@{domain}" - if "@" in normalized: return normalized - if _BARE_PHONE_RE.fullmatch(normalized): digits = re.sub(r"\D+", "", normalized) if digits: return f"{digits}@s.whatsapp.net" - return normalized def expand_whatsapp_aliases(identifier: str) -> Set[str]: """Return all identifiers transitively reachable via the bridge's ``lid-mapping-*.json`` files. - - Always includes the normalized input itself, so callers can ``in``-check - without a fallback branch. Empty set if ``identifier`` normalizes to empty. + Always includes the normalized input itself (callers can ``in``-check without a fallback); + empty set if ``identifier`` normalizes to empty. """ normalized = normalize_whatsapp_identifier(identifier) if not normalized: return set() - session_dir = get_hermes_dir("platforms/whatsapp/session", "whatsapp/session") resolved: Set[str] = set() queue = [normalized] - while queue: current = queue.pop(0) - if not current or current in resolved: + # _SAFE_IDENTIFIER_RE: defense-in-depth against path separators / traversal in the + # ``lid-mapping-{current}`` filename (the fixed prefix already prevents escape). + if not current or current in resolved or not _SAFE_IDENTIFIER_RE.match(current): continue - # Defense-in-depth against path separators / traversal in the - # ``lid-mapping-{current}`` filename; the fixed prefix already prevents - # escape, but this avoids depending on that filesystem-layout invariant. - if not _SAFE_IDENTIFIER_RE.match(current): - continue - resolved.add(current) for suffix in ("", "_reverse"): mapping_path = session_dir / f"lid-mapping-{current}{suffix}.json" @@ -109,21 +86,13 @@ def expand_whatsapp_aliases(identifier: str) -> Set[str]: continue if mapped and mapped not in resolved: queue.append(mapped) - return resolved def canonical_whatsapp_identifier(identifier: str) -> str: - """Return a stable sender identity across phone-JID/LID variants. - - Applies to DM ``chat_id`` and group ``participant_id`` alike (the bridge - may flip between forms for the same human). Picks the shortest alias from - :func:`expand_whatsapp_aliases`, which degrades to the normalized input - when no mapping files exist yet. Empty string for empty input. + """Return a stable sender identity across phone-JID/LID variants (DM ``chat_id`` and group + ``participant_id`` alike): the shortest alias from :func:`expand_whatsapp_aliases`, which + degrades to the normalized input when no mapping files exist. ``""`` for empty input. """ - normalized = normalize_whatsapp_identifier(identifier) - if not normalized: - return "" - - aliases = expand_whatsapp_aliases(normalized) - return min(aliases, key=lambda candidate: (len(candidate), candidate)) + aliases = expand_whatsapp_aliases(identifier) + return min(aliases, key=lambda c: (len(c), c)) if aliases else ""