refactor(gateway): stream/runner — checkpoint: consumer init/reset helpers, sentinel dispatch table, fallback/transport dedupe, watcher/voice/lease compaction
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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))))
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 = ""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 ""
|
||||
|
||||
Reference in New Issue
Block a user