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