refactor(gateway): voice/watchers/tts/lease — phase helpers, folded probes, compact docstrings

This commit is contained in:
Teknium
2026-09-02 20:26:41 -07:00
parent ff26aa9247
commit 0c58b7a9e4
4 changed files with 265 additions and 291 deletions
+118 -123
View File
@@ -16,6 +16,8 @@ import re
import sys
import time
from contextlib import suppress
from difflib import SequenceMatcher
from types import SimpleNamespace
from typing import TYPE_CHECKING, Dict, List, Optional
from gateway.config import Platform
@@ -30,49 +32,47 @@ 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"
_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:
"""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.
"""
"""``<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 ""
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. Voice mode belongs to the (bot, chat) pair, so the
namespace is the profile that OWNS the receiving adapter, not the routed profile."""
"""Voice mode belongs to the (bot, chat) pair: 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"):
cb = functools.partial(self._handle_voice_channel_input, adapter=adapter)
adapter._voice_input_callback = cb
adapter._voice_input_callback = functools.partial(
self._handle_voice_channel_input, adapter=adapter
)
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 {}
result = {}
for chat_id, mode in (data.items() if isinstance(data, dict) else ()):
if mode not in {"off", "voice_only", "all"}:
continue
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.", str(chat_id),
)
return result
if not isinstance(data, dict):
return {}
items = {str(k): m for k, m in data.items() if m in _VOICE_MODES}
for chat_id 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.", chat_id,
)
return {k: m for k, m in items.items() if ":" in k}
def _save_voice_modes(self) -> None:
try:
@@ -84,20 +84,19 @@ 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).
"""
"""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 not on:
target.discard(chat_id)
else:
target.add(chat_id)
if isinstance(other, set):
other.discard(chat_id)
return
target.add(chat_id)
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."""
@@ -108,18 +107,17 @@ class GatewayVoiceMixin:
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: ``_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``.
"""
"""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``."""
platform = getattr(adapter, "platform", None)
if not isinstance(platform, Platform):
return
chat_sets = [
(getattr(adapter, name, None), modes)
(chats, modes)
for name, modes in ((_OFF_SET, {"off"}), (_ON_SET, {"voice_only", "all"}))
if isinstance(chats := getattr(adapter, name, None), set)
]
chat_sets = [(chats, modes) for chats, modes in chat_sets if isinstance(chats, set)]
if not chat_sets:
return
# Lazy import: no module-level dep from gateway -> hermes_cli.
@@ -144,9 +142,7 @@ class GatewayVoiceMixin:
raw = getattr(event, "raw_message", None)
if getattr(raw, "guild_id", None): # slash command interaction
return int(raw.guild_id)
if getattr(raw, "guild", None): # regular message
return raw.guild.id
return None
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."""
@@ -177,12 +173,10 @@ class GatewayVoiceMixin:
except Exception as e:
logger.warning("Failed to join voice channel: %s", e)
adapter._voice_input_callback = None
if any(tok in str(e).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 not any(tok in str(e).lower() for tok in ("pynacl", "nacl", "davey")):
return f"Failed to join voice channel: {e}"
return ("Voice dependencies are missing (PyNaCl / davey). "
f"Install with: `{sys.executable} -m pip install PyNaCl`")
if not success:
adapter._voice_input_callback = None
return "Failed to join voice channel. Check bot permissions (Connect + Speak)."
@@ -200,11 +194,9 @@ class GatewayVoiceMixin:
"""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")
or not hasattr(adapter, "is_in_voice_channel")
or not 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:
@@ -220,10 +212,7 @@ class GatewayVoiceMixin:
def _handle_voice_timeout_cleanup(self, chat_id: str, *, adapter=None) -> None:
"""Adapter callback on voice-channel timeout: clear runner-side voice_mode state.
``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`` (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)
@@ -239,8 +228,6 @@ class GatewayVoiceMixin:
"""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"[^\w\s]", "", re.sub(r"\s+", " ", transcript).strip().lower())
if not normalized:
return False
@@ -250,43 +237,53 @@ class GatewayVoiceMixin:
if not isinstance(recent_store, dict):
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 or (
if any(
prior == normalized or (
len(prior) >= 16 and len(normalized) >= 16
and SequenceMatcher(None, prior, normalized).ratio() >= 0.95
):
recent_store[key] = recent
return True
recent.append((now, normalized))
recent_store[key] = recent[-5:]
)
for _, prior in recent
):
recent_store[key] = recent
return True
recent_store[key] = (recent + [(now, normalized)])[-5:]
return False
@staticmethod
def _voice_input_source(adapter, guild_id: int, user_id: int, text_ch_id) -> SessionSource:
"""Bound text channel's own source when available (voice shares the text conversation's
session), else a synthetic one."""
if source_data := getattr(adapter, "_voice_sources", {}).get(guild_id):
source = SessionSource.from_dict(source_data)
source.user_id = source.user_name = str(user_id)
return source
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
async def _handle_voice_channel_input(
self, guild_id: int, user_id: int, transcript: str, *, adapter=None
):
"""Handle transcribed voice from a user in a voice channel.
``adapter`` is the Discord adapter that captured the audio (bound via
``_bind_voice_input_callback``); under multiplexing each profile's bot must dispatch
through its own adapter, never the default profile's.
"""
"""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."""
if adapter is None:
adapter = self.adapters.get(Platform.DISCORD)
text_ch_id = adapter._voice_text_channels.get(guild_id) if adapter else None
if not text_ch_id:
return
# 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 = 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",
profile=getattr(adapter, "_owner_profile", None),
)
source = self._voice_input_source(adapter, guild_id, user_id, text_ch_id)
if not self._is_user_authorized(source):
logger.debug("Unauthorized voice input from user %d, ignoring", user_id)
return
@@ -302,34 +299,22 @@ class GatewayVoiceMixin:
if channel:
safe_text = transcript[:2000].replace("@everyone", "@\u200beveryone").replace("@here", "@\u200bhere")
await channel.send(f"**[Voice]** <@{user_id}>: {safe_text}")
# 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):
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.
from types import SimpleNamespace
event = MessageEvent(
source=source, text=transcript, message_type=MessageType.VOICE,
raw_message=SimpleNamespace(guild_id=guild_id, guild=None),
channel_prompt=channel_prompt,
channel_prompt=self._voice_channel_prompt(adapter, text_ch_id),
)
await adapter.handle_message(event)
def _should_send_voice_reply(
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 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.
"""
"""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."""
if not response or response.startswith("Error:"):
return False
chat_id = event.source.chat_id
@@ -353,14 +338,13 @@ class GatewayVoiceMixin:
)
return False
# 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
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 turn_messages if msg.get("role") == "assistant"
for msg in agent_messages[start:] if msg.get("role") == "assistant"
for tc in (msg.get("tool_calls") or [])
):
return False
@@ -373,12 +357,39 @@ class GatewayVoiceMixin:
"""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] = []
try:
from tools.tts_tool import text_to_speech_tool, _strip_markdown_for_tts
from tools.tts_tool import _strip_markdown_for_tts
tts_text = _strip_markdown_for_tts(text)
if not tts_text:
@@ -386,24 +397,9 @@ class GatewayVoiceMixin:
# 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
)
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
# One combined file or several separately valid files (combination unavailable or
# over a platform limit); legacy single-file results keep working.
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
await self._deliver_voice_reply(event, actual_paths)
actual_paths = await self._synthesize_voice_reply(tts_text, audio_path)
if actual_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:
@@ -417,12 +413,11 @@ class GatewayVoiceMixin:
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):
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)
+76 -80
View File
@@ -41,15 +41,13 @@ 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 (``on_session_finalize`` hooks, cached
agent teardown, cache eviction, ``expiry_finalized`` flag) and runs the cache/store sweeps.
"""
"""Finalize expired sessions (``on_session_finalize`` hooks, cached agent teardown, cache
eviction, ``expiry_finalized`` flag) and 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:
expired = await self._collect_expired_sessions()
if expired:
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)",
@@ -70,24 +68,20 @@ class GatewaySessionWatchersMixin:
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
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.
"""
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:
@@ -103,7 +97,20 @@ class GatewaySessionWatchersMixin:
)
store = self.async_session_store
await store.set_expiry_finalized(entry, clear_model_override=False)
failures.pop(sid, None)
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
async def _finalize_expired_session(self, key: str, entry) -> None:
"""Run finalize hooks, tear down the cached agent, clear conversation scope, persist."""
@@ -119,17 +126,7 @@ class GatewaySessionWatchersMixin:
)
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
agent = self._agent_for_expired_session(key)
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
@@ -177,7 +174,7 @@ class GatewaySessionWatchersMixin:
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()]
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:
@@ -185,9 +182,8 @@ class GatewaySessionWatchersMixin:
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()``; no turn-start or pending-inbound clocks.
"""
"""Activity snapshot for stall progress: the single source is
``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)
if agent is None or agent is _AGENT_PENDING_SENTINEL:
@@ -226,8 +222,7 @@ class GatewaySessionWatchersMixin:
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).
"""
notifications sent this pass (for tests)."""
from gateway.session_stall import (
resolve_session_idle_seconds_from_activity,
should_clear_session_stall_notification,
@@ -237,9 +232,7 @@ class GatewaySessionWatchersMixin:
notified_map = getattr(self, "_session_stall_notified", None)
if notified_map is None:
notified_map = self._session_stall_notified = {}
sent = 0
now = time.time()
candidates = self._stall_candidates()
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)
@@ -266,46 +259,11 @@ class GatewaySessionWatchersMixin:
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."""
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,
resolve_session_idle_seconds_from_activity,
)
from gateway.session_stall import format_session_stall_notification
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)
@@ -313,7 +271,7 @@ class GatewaySessionWatchersMixin:
# 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),
adapter.send(str(source.chat_id), notice, metadata=metadata),
timeout=_STALL_NOTIFY_SEND_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
@@ -332,13 +290,52 @@ class GatewaySessionWatchersMixin:
except Exception as exc:
logger.warning("Session stall notify failed for %s: %s", session_key, exc)
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:
"""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
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 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
if not await self._send_stall_notice(session_key, adapter, source, idle_seconds):
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.
"""
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
@@ -364,8 +361,7 @@ class GatewaySessionWatchersMixin:
await asyncio.sleep(min(30.0, max(1.0, float(interval))))
while self._running:
try:
timeout = self._session_stall_timeout_seconds()
if timeout > 0:
if (timeout := self._session_stall_timeout_seconds()) > 0:
await self._check_session_stalls(timeout)
except Exception as exc:
logger.debug("Session stall watcher error: %s", exc)
+54 -60
View File
@@ -1,12 +1,10 @@
"""Gateway streaming-TTS consumer — LLM deltas to adapter PCM audio sink.
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).
``on_delta`` (agent worker thread) never blocks: SentenceChunker -> thread-safe queue. The
``_run`` task on the gateway loop drains, synthesises via a ``StreamingTTSProvider`` and writes
PCM so playback starts mid-generation. Outcome: success -> ``completed``; failure before 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
@@ -41,10 +39,7 @@ class StreamingTTSConsumer:
) -> None:
from tools.tts_streaming import SentenceChunker, resolve_streaming_provider
self._adapter = adapter
self._chat_id = chat_id
self._loop = loop
self._metadata = metadata
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()
@@ -64,20 +59,13 @@ class StreamingTTSConsumer:
self._lock = threading.Lock()
self._strip_markdown = None # lazily imported to avoid import cycles
# 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())
active = property(lambda self: self._streamer is not None) # usable streaming provider
completed = property(lambda self: self._completed) # streaming audio fully delivered
partial = property(lambda self: self._partial) # some audio audible before a failure/drop
audible = property(lambda self: bool(self._handle and self._handle.audible)) # PCM written
dropped = property(lambda self: self._dropped) # queue saturation dropped >= 1 clause
suppress_whole_file = property(lambda self: self._suppress_whole_file) # skip whole-file TTS
done = property(lambda self: self._task is not None and self._task.done()) # drain task ended
def _enqueue_clauses(self, clauses, full_msg: str, *, log_errors: bool) -> None:
try:
@@ -100,8 +88,7 @@ class StreamingTTSConsumer:
def finish(self) -> None:
"""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``.
"""
clauses so the drain loop ends deterministically without racing a late ``on_delta``."""
if self._finished:
return
self._finished = True
@@ -121,13 +108,13 @@ class StreamingTTSConsumer:
self._queue.put_nowait(sentinel)
return True
except queue.Full:
try:
self._queue.get_nowait()
if mark_dropped:
self._dropped = True
except queue.Empty:
return not mark_dropped
return False
pass
try:
self._queue.get_nowait()
except queue.Empty:
return not mark_dropped
self._dropped = self._dropped or mark_dropped
return False
def start(self) -> asyncio.Task:
"""Create (once) and return the async drain task on the gateway loop."""
@@ -138,20 +125,18 @@ class StreamingTTSConsumer:
def _settle(self, *, failed: bool) -> None:
"""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
audible, degraded = self._handle.audible, failed or self._dropped
self._completed = audible and not degraded
if audible and degraded:
self._partial = True
self._partial = self._partial or (audible and degraded)
self._suppress_whole_file = audible
async def _run(self) -> None:
"""Drain clauses from the queue, synthesise, and write to the adapter."""
async def _open_handle(self) -> bool:
"""Open the adapter's streaming-audio handle; False when unsupported or begin failed."""
if not self.active:
return
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", "?"))
return
return False
try:
self._handle = await self._adapter.begin_streaming_tts(
self._chat_id, self._audio_format, metadata=self._metadata
@@ -159,27 +144,36 @@ class StreamingTTSConsumer:
except Exception as exc:
logger.debug("begin_streaming_tts failed: %s", exc)
self._handle = None
return
if self._handle is 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."""
if not await self._open_handle():
return
self._suppress_whole_file = False
try:
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 await self._drain():
return
if not self._aborted and self._handle is not None:
try:
await self._adapter.finish_streaming_tts(self._handle, interrupted=self._aborted)
+17 -28
View File
@@ -3,13 +3,12 @@
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`.
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`.
"""
import asyncio
@@ -58,10 +57,8 @@ class TurnLeaseToken:
self.released = False
def __repr__(self) -> str: # pragma: no cover - debug aid
return (
f"TurnLeaseToken(session_id={self.session_id!r}, owner_key={self.owner_key!r}, "
f"generation={self.generation}, released={self.released})"
)
return (f"TurnLeaseToken(session_id={self.session_id!r}, owner_key={self.owner_key!r}, "
f"generation={self.generation}, released={self.released})")
class _SessionLease:
@@ -100,16 +97,12 @@ class SessionTurnLeaseRegistry:
return lease
def _evict_idle(self) -> None:
"""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:
"""Drop oldest idle entries to fit a new lease under the cap; never a held/contended one."""
if (overflow := len(self._leases) - self._max_entries + 1) <= 0:
return
idle_ids = sorted(
(sid for sid, lease in self._leases.items() if lease.idle),
key=lambda sid: self._leases[sid].last_used,
)
for sid in idle_ids[:overflow]:
idle = sorted((sid for sid, l in self._leases.items() if l.idle),
key=lambda sid: self._leases[sid].last_used)
for sid in idle[:overflow]:
self._leases.pop(sid, None)
async def acquire(
@@ -158,15 +151,11 @@ class SessionTurnLeaseRegistry:
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) — 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).
"""
(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
token is None or token.released or not new_session_id
or new_session_id == token.session_id
):
return False