refactor(gateway): stream/runner — checkpoint: consumer init/reset helpers, sentinel dispatch table, fallback/transport dedupe, watcher/voice/lease compaction

This commit is contained in:
Teknium
2026-09-02 19:19:38 -07:00
parent 113f04616b
commit 6ad463a3ca
14 changed files with 1221 additions and 1974 deletions

View File

@@ -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>:<platform>:<chat_id>`` (profile whose bot speaks);
the default profile keeps ``<platform>:<chat_id>`` 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: ``<profile>:<platform>:<chat_id>`` 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 ``<platform>:<chat_id>`` 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,
)

View File

@@ -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))))

View File

@@ -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.^=)]"

File diff suppressed because it is too large Load Diff

View File

@@ -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)

View File

@@ -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

View File

@@ -25,6 +25,34 @@ class StreamThinkFilterMixin:
"</THINKING>", "</thinking>", "</thought>",
)
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 <Think>, <THINKING>, …
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 </think> (thinking-mode toggle dropped the
# open, or incomplete upstream stripping) is noise.
# An orphan </think> (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] == "</":
for tag in cls._CLOSE_THINK_TAGS:
tag_lower = tag.lower()
tag_len = len(tag_lower)
if text_lower[i:i + tag_len] == tag_lower:
j = i + tag_len
if text_lower.startswith(tag_lower, i):
j = i + len(tag_lower)
while j < len(text) and text[j] in " \t\n\r":
j += 1
i = j

View File

@@ -22,13 +22,7 @@ class StreamTransportMixin:
_MIN_NEW_MSG_CHARS = 4
async def _edit_message(
self,
*,
message_id: str,
content: str,
finalize: bool = False,
):
async def _edit_message(self, *, message_id: str, content: str, finalize: bool = False):
"""Edit via the adapter, passing routing metadata when supported."""
# Contract: adapters must accept finalize= even when False (test-guarded).
kwargs = {
@@ -41,8 +35,7 @@ class StreamTransportMixin:
try:
params = inspect.signature(self.adapter.edit_message).parameters
if "metadata" in params or any(
param.kind is inspect.Parameter.VAR_KEYWORD
for param in params.values()
param.kind is inspect.Parameter.VAR_KEYWORD for param in params.values()
):
kwargs["metadata"] = self.metadata
except (TypeError, ValueError):
@@ -52,10 +45,7 @@ class StreamTransportMixin:
async def _send_seed_frame(self):
"""Open a native stream with an empty seed frame (typing indicator before any token)."""
return await self.adapter.send_stream_frame(
"",
chat_id=self.chat_id,
reply_to=self._initial_reply_to_id,
turn_id=self._turn_id,
"", chat_id=self.chat_id, reply_to=self._initial_reply_to_id, turn_id=self._turn_id,
)
async def _send_frame(self, text: str, *, finalize: bool):
@@ -73,6 +63,15 @@ class StreamTransportMixin:
self._native_stream_opened = False
self._native_last_pushed_len = 0
async def _close_empty_native_bubble(self, fail_log: str) -> 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

View File

@@ -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:

View File

@@ -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 = ""

View File

@@ -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)

View File

@@ -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

View File

@@ -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

View File

@@ -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
``<digits>@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.
``<digits>@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 ""