refactor(gateway): voice/watchers/tts/lease — phase helpers, folded probes, compact docstrings
This commit is contained in:
+118
-123
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user