diff --git a/gateway/session.py b/gateway/session.py index 1c536e6226..24e688de70 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -1,11 +1,7 @@ """ -Session management for the gateway. - -Handles: -- Session context tracking (where messages come from) -- Session storage (conversations persisted to disk) -- Reset policy evaluation (when to start fresh) -- Dynamic system prompt injection (agent knows its context) +Session management for the gateway: source tracking (where messages come +from), the persisted routing index (SessionStore), reset policy evaluation and +the dynamic "Current Session Context" system prompt section. """ import asyncio @@ -14,67 +10,36 @@ import logging import os import json import threading -import uuid from pathlib import Path from datetime import datetime, timedelta from dataclasses import dataclass, field, fields from typing import Dict, List, Optional, Any +from .config import ( + Platform, + GatewayConfig, + SessionResetPolicy, # noqa: F401 — re-exported via gateway/__init__.py + HomeChannel, +) +from .whatsapp_identity import ( + canonical_whatsapp_identifier, + normalize_whatsapp_identifier, # noqa: F401 - re-exported for gateway.session callers +) +from gateway.session_persistence import SessionPersistenceMixin, _DB_UNPINNED # noqa: F401 +from gateway.session_recovery import SessionRecoveryMixin +from gateway.session_lifecycle import ( # noqa: F401 — _now & co. re-exported for callers/tests + SessionLifecycleMixin, + _iso, + _new_session_id, + _now, + _parse_iso, + auto_continue_freshness_window, +) +from gateway.session_transcript import SessionTranscriptMixin, TranscriptReadError # noqa: F401 + logger = logging.getLogger(__name__) -class TranscriptReadError(RuntimeError): - """Raised when persisted history cannot be read safely.""" - - def __init__(self, session_id: str) -> None: - self.session_id = session_id - super().__init__(f"transcript read failed for session {session_id}") - - -def _now() -> datetime: - """Return the current local time.""" - return datetime.now() - - -def _new_session_id(now: datetime) -> str: - return f"{now.strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:8]}" - - -def _iso(dt: Optional[datetime]) -> Optional[str]: - return dt.isoformat() if dt else None - - -def _parse_iso(value) -> Optional[datetime]: - """``datetime.fromisoformat`` that returns None for empty/malformed input.""" - if not value: - return None - try: - return datetime.fromisoformat(value) - except (TypeError, ValueError): - return None - - -# Default auto-continue freshness window (1 hour): a restart-interrupted -# session is only auto-resumed while within this window of when -# ``resume_pending`` was marked. ``gateway/run.py`` bridges config.yaml -# ``agent.gateway_auto_continue_freshness`` into the env var at startup. -_AUTO_CONTINUE_FRESHNESS_SECS_DEFAULT = 60 * 60 - - -def auto_continue_freshness_window() -> float: - """Auto-continue freshness window in seconds (single source of truth for - the resume scheduler and the routing-time zombie gate). - - Reads ``HERMES_AUTO_CONTINUE_FRESHNESS``; falls back to the default when - unset or malformed. Non-positive disables the gate. - """ - raw = os.environ.get("HERMES_AUTO_CONTINUE_FRESHNESS") - try: - return float(raw) if raw else float(_AUTO_CONTINUE_FRESHNESS_SECS_DEFAULT) - except (TypeError, ValueError): - return float(_AUTO_CONTINUE_FRESHNESS_SECS_DEFAULT) - - # --------------------------------------------------------------------------- # PII redaction helpers # --------------------------------------------------------------------------- @@ -90,50 +55,27 @@ def _hash_sender_id(value: str) -> str: def _hash_chat_id(value: str) -> str: - """Hash the numeric portion of a chat ID, preserving platform prefix. - - ``telegram:12345`` → ``telegram:`` - ``12345`` → ```` - """ + """Hash the numeric portion of a chat ID, preserving a ``platform:`` prefix.""" colon = value.find(":") if colon > 0: - prefix = value[:colon] - return f"{prefix}:{_hash_id(value[colon + 1:])}" + return f"{value[:colon]}:{_hash_id(value[colon + 1:])}" return _hash_id(value) -from .config import ( - Platform, - GatewayConfig, - SessionResetPolicy, # noqa: F401 — re-exported via gateway/__init__.py - HomeChannel, -) -from .whatsapp_identity import ( - canonical_whatsapp_identifier, - normalize_whatsapp_identifier, # noqa: F401 - re-exported for gateway.session callers -) -from gateway.session_persistence import SessionPersistenceMixin -from gateway.session_recovery import SessionRecoveryMixin -from gateway.session_lifecycle import SessionLifecycleMixin -from gateway.session_transcript import SessionTranscriptMixin - def _is_path_unsafe(value: object, *, strict: bool = True) -> bool: """True if ``value`` could traverse outside the sessions dir. - Session ids become filenames (``sessions_dir / f"{session_id}.json"``), so - the strict form rejects ``..``, ANY path separator, and a leading Windows - drive letter. The relaxed form (``strict=False``) is for *logical* session - keys, where interior ``/`` is legitimate (Google Chat + Session ids become filenames, so the strict form rejects ``..``, ANY path + separator, and a leading Windows drive letter. ``strict=False`` is for + *logical* session keys, where interior ``/`` is legitimate (Google Chat ``spaces//threads/``): only a *leading* separator is rejected. """ if not value: return False s = str(value) - if ".." in s: + if ".." in s or (strict and ("/" in s or "\\" in s)): return True - if strict and ("/" in s or "\\" in s): - return True - if not strict and (s.startswith("/") or s.startswith("\\")): + if not strict and s.startswith(("/", "\\")): return True return len(s) >= 2 and s[0].isalpha() and s[1] == ":" @@ -151,43 +93,39 @@ class SessionSource: chat_type: str = "dm" # "dm", "group", "channel", "thread" user_id: Optional[str] = None user_name: Optional[str] = None - thread_id: Optional[str] = None # For forum topics, Discord threads, etc. - chat_topic: Optional[str] = None # Channel topic/description (Discord, Slack) - user_id_alt: Optional[str] = None # Platform-specific stable alt ID (Signal UUID, Feishu union_id) + thread_id: Optional[str] = None # forum topics, Discord threads, etc. + chat_topic: Optional[str] = None # channel topic/description (Discord, Slack) + user_id_alt: Optional[str] = None # platform-specific stable alt ID (Signal UUID, Feishu union_id) chat_id_alt: Optional[str] = None # Signal group internal ID - is_bot: bool = False # True when the message author is a bot/webhook (Discord) + is_bot: bool = False # message author is a bot/webhook (Discord) # Platform-neutral SCOPE discriminator (Discord guild / Slack workspace / - # Matrix server); drives server/workspace isolation. `scope_id` is - # canonical; `guild_id` is a deprecated alias kept during the cross-repo - # dual-read/dual-write overlap (both written, scope_id wins on read). + # Matrix server) driving server/workspace isolation. ``scope_id`` is + # canonical; ``guild_id`` is a deprecated alias kept while the cross-repo + # dual-read/dual-write overlap lasts (both written, scope_id wins on read). scope_id: Optional[str] = None - guild_id: Optional[str] = None # @deprecated legacy alias for scope_id - parent_chat_id: Optional[str] = None # Parent channel when chat_id refers to a thread - message_id: Optional[str] = None # ID of the triggering message (for pin/reply/react) - role_authorized: bool = False # True when adapter granted access via role (not user ID) + guild_id: Optional[str] = None + parent_chat_id: Optional[str] = None # parent channel when chat_id is a thread + message_id: Optional[str] = None # triggering message (pin/reply/react) + role_authorized: bool = False # adapter granted access via role, not user ID # Profile this message is routed to in a multiplexing gateway (None => - # active/default). Drives session-key namespacing and per-turn scope. + # active/default); drives session-key namespacing and per-turn scope. profile: Optional[str] = None # Transport-local fail-closed signal for an explicit profile route whose # target is not served. Excluded from repr/equality and wire serialization. profile_route_rejected: bool = field(default=False, repr=False, compare=False) - # Discord auto-thread metadata: explicit so pre-existing or human-renamed # threads are never mistaken for safe rename targets. auto_thread_created: bool = False auto_thread_initial_name: Optional[str] = None - # Discord auto-thread continuity: set by the connector on a CHANNEL message - # (no thread_id yet) that WILL be delivered into a new thread whose id == - # this message id. Keying the session on it makes the initiating channel - # message and later in-thread follow-ups share ONE session. + # that WILL be delivered into a new thread whose id == this message id, so + # the initiating message and later in-thread follow-ups share ONE session. prospective_thread_id: Optional[str] = None - - # Wire-INVISIBLE trust signal: True when delivered over the per-instance - # authenticated relay WebSocket, whose connector already resolved - # owner-only author bindings. ``platform`` carries the UNDERLYING platform - # (not ``relay``), so authz must key upstream trust off THIS flag. - # Excluded from to_dict/from_dict so a peer can never forge or persist it. + # Wire-INVISIBLE trust signal: delivered over the per-instance authenticated + # relay WebSocket whose connector already resolved owner-only author + # bindings. ``platform`` carries the UNDERLYING platform (not ``relay``), so + # authz must key upstream trust off THIS flag. Excluded from to_dict/ + # from_dict so a peer can never forge or persist it. delivered_via_upstream_relay: bool = False def __post_init__(self) -> None: @@ -202,22 +140,17 @@ class SessionSource: def _describe(chat_type: str, user_label: str, chat_label: str) -> str: if chat_type == "dm": return f"DM with {user_label}" - prefix = _CHAT_TYPE_PREFIX.get(chat_type, "") - return f"{prefix}{chat_label}" + return f"{_CHAT_TYPE_PREFIX.get(chat_type, '')}{chat_label}" @property def description(self) -> str: """Human-readable description of the source.""" if self.platform == Platform.LOCAL: return "CLI terminal" - parts = [self._describe( - self.chat_type, - self.user_name or self.user_id or "user", - self.chat_name or self.chat_id, - )] - if self.thread_id: - parts.append(f"thread: {self.thread_id}") - return ", ".join(parts) + desc = self._describe( + self.chat_type, self.user_name or self.user_id or "user", self.chat_name or self.chat_id + ) + return f"{desc}, thread: {self.thread_id}" if self.thread_id else desc # Wire layout (order matters for byte-stable JSON): always-present fields, # then truthy-only optionals around the dual-written scope pair. @@ -231,17 +164,13 @@ class SessionSource: d.update((name, getattr(self, name)) for name in self._ALWAYS_FIELDS) def _optional(names) -> None: - for name in names: - value = getattr(self, name) - if value: - d[name] = value + d.update((name, v) for name in names if (v := getattr(self, name))) _optional(self._OPTIONAL_PRE_SCOPE) # Dual-write scope_id + deprecated guild_id alias during the migration. scope = self.scope_id if self.scope_id is not None else self.guild_id if scope: - d["scope_id"] = scope - d["guild_id"] = scope + d["scope_id"] = d["guild_id"] = scope _optional(self._OPTIONAL_POST_SCOPE) if self.auto_thread_created: d["auto_thread_created"] = True @@ -265,7 +194,6 @@ class SessionSource: ) - @dataclass class SessionContext: """Full session context for dynamic system prompt injection.""" @@ -273,8 +201,6 @@ class SessionContext: connected_platforms: List[Platform] home_channels: Dict[Platform, HomeChannel] shared_multi_user_session: bool = False - - # Session metadata session_key: str = "" session_id: str = "" created_at: Optional[datetime] = None @@ -284,9 +210,7 @@ class SessionContext: return { "source": self.source.to_dict(), "connected_platforms": [p.value for p in self.connected_platforms], - "home_channels": { - p.value: hc.to_dict() for p, hc in self.home_channels.items() - }, + "home_channels": {p.value: hc.to_dict() for p, hc in self.home_channels.items()}, "shared_multi_user_session": self.shared_multi_user_session, "session_key": self.session_key, "session_id": self.session_id, @@ -295,25 +219,22 @@ class SessionContext: } +# Platforms where user IDs can be redacted: no ``<@user_id>``-style mention +# system that needs raw IDs (which is why Discord is excluded). _PII_SAFE_PLATFORMS = frozenset({ - Platform.WHATSAPP, - Platform.SIGNAL, - Platform.TELEGRAM, - Platform.BLUEBUBBLES, + Platform.WHATSAPP, Platform.SIGNAL, Platform.TELEGRAM, Platform.BLUEBUBBLES, }) -"""Platforms where user IDs can be redacted (no ``<@user_id>``-style mention -system that needs raw IDs — which is why Discord is excluded).""" def _slack_tools_loaded() -> bool: """True iff the agent will actually have Slack tools this session. - Either (1) the native `slack` toolset is enabled for the platform AND - `SLACK_BOT_TOKEN` is set (the tool's `check_fn` gates on it), or (2) an - MCP server whose name suggests Slack has ACTUALLY registered tools into - the live registry (configured-but-unconnected servers don't count; MCP - servers are process-wide, so this is intentionally not per-session). - Returns False on any error so a bad config never promises missing tools. + Either the native `slack` toolset is enabled for the platform AND + `SLACK_BOT_TOKEN` is set (the tool's `check_fn` gates on it), or an MCP + server whose name suggests Slack has ACTUALLY registered tools (configured + -but-unconnected servers don't count; MCP servers are process-wide, so this + is intentionally not per-session). False on any error so a bad config + never promises missing tools. """ try: from tools.mcp_tool import get_registered_mcp_server_names @@ -376,8 +297,8 @@ def neutralize_untrusted_inline_text(value: Any, *, max_chars: int = _MAX_PROMPT Sibling of :func:`_format_untrusted_prompt_value` for inline call sites (e.g. a ``[Name]`` turn prefix) where JSON-quoting would visibly change - rendering. Embedded newlines are the injection vector: they let a display - name masquerade as a new markdown section. Collapsing them keeps a normal + rendering. Embedded newlines are the injection vector (a display name + masquerading as a new markdown section); collapsing them keeps a normal value byte-identical while making a hostile one inert. """ text = str(value).replace("\r\n", "\n").replace("\r", "\n").replace("\n", " ") @@ -494,11 +415,12 @@ def build_session_context_prompt( ``pii_safe``), user/chat IDs are replaced with deterministic hashes for the LLM only; routing keeps the originals in SessionSource. """ - _is_pii_safe = context.source.platform in _PII_SAFE_PLATFORMS + src = context.source + _is_pii_safe = src.platform in _PII_SAFE_PLATFORMS if not _is_pii_safe: try: from gateway.platform_registry import platform_registry - entry = platform_registry.get(context.source.platform.value) + entry = platform_registry.get(src.platform.value) if entry and entry.pii_safe: _is_pii_safe = True except Exception: @@ -519,12 +441,10 @@ def build_session_context_prompt( "", ] - # Source info - platform_name = context.source.platform.value.title() - if context.source.platform == Platform.LOCAL: + platform_name = src.platform.value.title() + if src.platform == Platform.LOCAL: lines.append(f"**Source:** {platform_name} (the machine running this agent)") else: - src = context.source if redact_pii: # Safe description without raw IDs (note: no thread suffix). desc = SessionSource._describe( @@ -534,17 +454,12 @@ def build_session_context_prompt( ) else: desc = src.description - lines.append( - f"**Source:** {platform_name} ({_format_untrusted_prompt_value(desc)})" - ) + lines.append(f"**Source:** {platform_name} ({_format_untrusted_prompt_value(desc)})") - if context.source.chat_topic: - lines.append( - f"**Channel Topic:** {_format_untrusted_prompt_value(context.source.chat_topic)}" - ) + if src.chat_topic: + lines.append(f"**Channel Topic:** {_format_untrusted_prompt_value(src.chat_topic)}") - if context.source.platform == Platform.MATRIX: - src = context.source + if src.platform == Platform.MATRIX: lines += [ "", f"**Matrix Room:** {_format_untrusted_prompt_value(src.chat_name or src.chat_id)}", @@ -562,28 +477,22 @@ def build_session_context_prompt( # prompt (changes per turn -> busts the prompt cache); sender names are # prefixed on each user message instead. if context.shared_multi_user_session: - session_label = "Multi-user thread" if context.source.thread_id else "Multi-user session" + session_label = "Multi-user thread" if src.thread_id else "Multi-user session" lines.append( f"**Session type:** {session_label} — messages are prefixed " "with [sender name]. Multiple users may participate." ) - elif context.source.user_name: - lines.append( - f"**User:** {_format_untrusted_prompt_value(context.source.user_name)}" - ) - elif context.source.user_id: - uid = context.source.user_id - if redact_pii: - uid = _hash_sender_id(uid) + elif src.user_name: + lines.append(f"**User:** {_format_untrusted_prompt_value(src.user_name)}") + elif src.user_id: + uid = _hash_sender_id(src.user_id) if redact_pii else src.user_id lines.append(f"**User ID:** {_format_untrusted_prompt_value(uid)}") - lines.extend(_PLATFORM_NOTES.get(context.source.platform, lambda ctx: [])(context)) - - platforms_list = ["local (files on this machine)"] - for p in context.connected_platforms: - if p != Platform.LOCAL: - platforms_list.append(f"{p.value}: Connected ✓") + lines.extend(_PLATFORM_NOTES.get(src.platform, lambda ctx: [])(context)) + platforms_list = ["local (files on this machine)"] + [ + f"{p.value}: Connected ✓" for p in context.connected_platforms if p != Platform.LOCAL + ] lines.append(f"**Connected Platforms:** {', '.join(platforms_list)}") if context.home_channels: @@ -597,17 +506,13 @@ def build_session_context_prompt( from hermes_constants import display_hermes_home - if context.source.platform == Platform.LOCAL: + if src.platform == Platform.LOCAL: lines.append("- `\"origin\"` → Local output (saved to files)") else: - _origin_label = _format_untrusted_prompt_value( - context.source.chat_name or _chat_label(context.source.chat_id) - ) + _origin_label = _format_untrusted_prompt_value(src.chat_name or _chat_label(src.chat_id)) lines.append(f"- `\"origin\"` → Back to this chat ({_origin_label})") - lines.append( - f"- `\"local\"` → Save to local files only ({display_hermes_home()}/cron/output/)" - ) + lines.append(f"- `\"local\"` → Save to local files only ({display_hermes_home()}/cron/output/)") for platform, home in context.home_channels.items(): home_name = _format_untrusted_prompt_value(home.name) lines.append(f"- `\"{platform.value}\"` → Home channel ({home_name})") @@ -641,19 +546,12 @@ class SessionEntry: session_id: str created_at: datetime updated_at: datetime - - # Origin metadata for delivery routing - origin: Optional[SessionSource] = None - - # Display metadata + origin: Optional[SessionSource] = None # delivery routing display_name: Optional[str] = None platform: Optional[Platform] = None chat_type: str = "dm" - - # Small, JSON-serializable per-entry state (e.g. Slack thread watermarks); - # persisted in the routing index. + # Small, JSON-serializable per-entry state (e.g. Slack thread watermarks). metadata: Dict[str, Any] = field(default_factory=dict) - # Token tracking input_tokens: int = 0 output_tokens: int = 0 @@ -662,32 +560,23 @@ class SessionEntry: total_tokens: int = 0 estimated_cost_usd: float = 0.0 cost_status: str = "unknown" - - # Last API-reported prompt tokens (for accurate compression pre-check) - last_prompt_tokens: int = 0 - + last_prompt_tokens: int = 0 # last API-reported prompt tokens (compression pre-check) # Set when a session was created because the previous one expired; - # consumed once by the message handler to inject a notice into context + # consumed once by the message handler to inject a notice into context. was_auto_reset: bool = False auto_reset_reason: Optional[str] = None # "idle" or "daily" - reset_had_activity: bool = False # whether the expired session had any messages - - # session_id replaced by an auto-reset; feeds build_channel_continuity_note. - prev_session_id: Optional[str] = None - + reset_had_activity: bool = False # the expired session had messages + prev_session_id: Optional[str] = None # replaced by an auto-reset; feeds build_channel_continuity_note # Set by reset_session() on explicit /new or /reset; consumed once to # re-inject topic/channel skills. Distinct from was_auto_reset, which # fires the "expired due to inactivity" notice (wrong for a manual reset). is_fresh_reset: bool = False - - # Set by the expiry watcher after finalizing an expired session; persisted - # so restarts don't re-run finalization. + # Set by the expiry watcher after finalizing; persisted so restarts don't + # re-run finalization. expiry_finalized: bool = False - # Next get_or_create_session() auto-resets (new session_id). Set by /stop # to break stuck-resume loops. suspended: bool = False - # Interrupted by a restart/drain timeout but recovery expected. Unlike # ``suspended``, the session_id is preserved so the agent auto-continues # the same transcript. Cleared after the next successful turn; escalation @@ -695,13 +584,11 @@ class SessionEntry: resume_pending: bool = False resume_reason: Optional[str] = None # e.g. "restart_timeout" last_resume_marked_at: Optional[datetime] = None - # Durable marker of the executing agent turn; CAS-cleared on normal # unwind, left behind by SIGKILL/OOM so unclean startup recovers the exact # interrupted session instead of guessing from ``updated_at``. active_turn_token: Optional[str] = None active_turn_started_at: Optional[datetime] = None - # Session-scoped /model override (model/provider/base_url ONLY — never # credentials; see sanitize_model_override). Persisted so a restart does # not revert sessions to the global default model. @@ -714,6 +601,9 @@ class SessionEntry: "total_tokens", "last_prompt_tokens", "estimated_cost_usd", "cost_status", "expiry_finalized", "suspended", "resume_pending", "resume_reason", ) + _RESET_FIELDS = ( + "is_fresh_reset", "was_auto_reset", "auto_reset_reason", "reset_had_activity", "prev_session_id", + ) def to_dict(self) -> Dict[str, Any]: result = { @@ -726,14 +616,11 @@ class SessionEntry: "chat_type": self.chat_type, "metadata": self.metadata, } - for name in self._PLAIN_FIELDS: - result[name] = getattr(self, name) + result.update((name, getattr(self, name)) for name in self._PLAIN_FIELDS) result["last_resume_marked_at"] = _iso(self.last_resume_marked_at) result["active_turn_token"] = self.active_turn_token result["active_turn_started_at"] = _iso(self.active_turn_started_at) - for name in ("is_fresh_reset", "was_auto_reset", "auto_reset_reason", - "reset_had_activity", "prev_session_id"): - result[name] = getattr(self, name) + result.update((name, getattr(self, name)) for name in self._RESET_FIELDS) if self.model_override: # Defence-in-depth against an unsanitized dict stored directly. result["model_override"] = sanitize_model_override(self.model_override) @@ -751,18 +638,15 @@ class SessionEntry: except ValueError as e: logger.debug("Unknown platform value %r: %s", data["platform"], e) - last_resume_marked_at = _parse_iso(data.get("last_resume_marked_at")) active_turn_started_at = _parse_iso(data.get("active_turn_started_at")) active_turn_token = data.get("active_turn_token") if not isinstance(active_turn_token, str) or not active_turn_token: # The token/timestamp pair is written atomically; a partial or # malformed pair is not trustworthy enough to auto-resume. - active_turn_token = None - active_turn_started_at = None + active_turn_token = active_turn_started_at = None session_key = data["session_key"] session_id = data["session_id"] - # CWE-22: session_id becomes a filename (strict guard); session_key is # a logical routing key where interior ``/`` is legitimate (relaxed). if _is_path_unsafe(session_id): @@ -771,7 +655,7 @@ class SessionEntry: raise ValueError("Invalid session_key: potential directory traversal detected") defaults = {f.name: f.default for f in fields(cls)} - plain = {name: data.get(name, defaults[name]) for name in cls._PLAIN_FIELDS} + plain = {name: data.get(name, defaults[name]) for name in cls._PLAIN_FIELDS + cls._RESET_FIELDS} plain["expiry_finalized"] = data.get("expiry_finalized", data.get("memory_flushed", False)) return cls( session_key=session_key, @@ -783,16 +667,11 @@ class SessionEntry: platform=platform, chat_type=data.get("chat_type", "dm"), metadata=dict(data.get("metadata") or {}), - **plain, - last_resume_marked_at=last_resume_marked_at, + last_resume_marked_at=_parse_iso(data.get("last_resume_marked_at")), active_turn_token=active_turn_token, active_turn_started_at=active_turn_started_at, - is_fresh_reset=data.get("is_fresh_reset", False), - was_auto_reset=data.get("was_auto_reset", False), - auto_reset_reason=data.get("auto_reset_reason"), - reset_had_activity=data.get("reset_had_activity", False), - prev_session_id=data.get("prev_session_id"), model_override=sanitize_model_override(data.get("model_override")), + **plain, ) @@ -804,9 +683,8 @@ def build_channel_continuity_note( After an auto-reset the agent could bind a new request to an unrelated recent session; this points it at the prior session in *this* channel so - it recalls context via ``session_search``. Returns ``None`` unless the - platform is Slack/Discord, the auto-reset had real activity, and the - previous session_id is recorded. + it recalls context via ``session_search``. ``None`` unless the platform is + Slack/Discord, the auto-reset had real activity, and prev_session_id is set. """ if source.platform not in (Platform.SLACK, Platform.DISCORD): return None @@ -852,6 +730,15 @@ def _session_key_namespace(profile: Optional[str]) -> str: return f"agent:{profile}" +def _canonical_participant(source: SessionSource) -> Optional[str]: + """Sender id for key isolation; WhatsApp JID/LID aliases are canonicalized + so alias flips cannot split one member into two sessions.""" + participant_id = source.user_id_alt or source.user_id + if participant_id and source.platform == Platform.WHATSAPP: + participant_id = canonical_whatsapp_identifier(str(participant_id)) or participant_id + return participant_id + + def build_session_key( source: SessionSource, group_sessions_per_user: bool = True, @@ -870,9 +757,7 @@ def build_session_key( ns = _session_key_namespace(profile) platform = source.platform.value slack_scope_id = ( - str(source.scope_id) - if source.platform == Platform.SLACK and source.scope_id - else None + str(source.scope_id) if source.platform == Platform.SLACK and source.scope_id else None ) if source.chat_type == "dm": dm_chat_id = source.chat_id @@ -887,22 +772,14 @@ def build_session_key( else: # No chat_id: fall back to the sender id before the bare # per-platform sink, or every chat_id-less DM shares one agent. - dm_participant_id = source.user_id_alt or source.user_id - if dm_participant_id and source.platform == Platform.WHATSAPP: - dm_participant_id = ( - canonical_whatsapp_identifier(str(dm_participant_id)) - or dm_participant_id - ) + dm_participant_id = _canonical_participant(source) if dm_participant_id: dm_parts.append(str(dm_participant_id)) if source.thread_id: dm_parts.append(source.thread_id) return ":".join(str(part) for part in dm_parts) - participant_id = source.user_id_alt or source.user_id - if participant_id and source.platform == Platform.WHATSAPP: - # JID/LID alias flips would otherwise split one member into two sessions. - participant_id = canonical_whatsapp_identifier(str(participant_id)) or participant_id + participant_id = _canonical_participant(source) # Discord auto-thread continuity: key a channel-initiating message on the # thread it WILL be delivered into (prospective_thread_id), and normalize # the chat_type slot to "thread" so in-thread follow-ups byte-match. A @@ -939,6 +816,31 @@ class _SessionFlight: self.error: Optional[BaseException] = None +@dataclass +class _RouteChecks: + """Lock-free I/O results for an existing route (phase 1b of a transition).""" + session_id: str # the entry's session_id when snapshotted + canonical_id: Optional[str] # compression tip (may equal session_id) + is_stale: bool # row already ended in state.db + reset_reason: Optional[str] + + +@dataclass +class _RouteDecision: + """What the locked apply-phase decided for one routing transition.""" + entry: Optional["SessionEntry"] = None + needs_save: bool = False + # Healthy-path saves take the single-row UPSERT fast path; structural + # transitions (recover/create) keep the full rewrite. + metadata_only_save: bool = False + needs_recover: bool = False + # Auto-reset bookkeeping: reason (None = no auto-reset), whether the ended + # session had activity, and its id (predecessor to end + continuity hint). + reset_reason: Optional[str] = None + reset_had_activity: bool = False + prev_session_id: Optional[str] = None + + class AsyncSessionStore: """Async boundary for the synchronous, thread-safe SessionStore.""" @@ -956,11 +858,6 @@ class AsyncSessionStore: return _offloaded -# "No SessionDB pinned" sentinel: lets ``_db`` distinguish "resolve from the -# active scope" from a deliberate ``store._db = None`` (JSONL fallback). -_DB_UNPINNED = object() - - class SessionStore( SessionPersistenceMixin, SessionRecoveryMixin, @@ -1053,9 +950,8 @@ class SessionStore( setattr(self, name, value) return value - def _has_active_processes_safe(self, session_key: str, *, context: str) -> bool: - """Return whether a session has active work, failing closed on registry errors.""" + """Whether a session has active work, failing closed on registry errors.""" if self._has_active_processes_fn is None: return False try: @@ -1063,23 +959,10 @@ class SessionStore( except Exception as exc: logger.warning( "has_active_processes_fn raised during %s for %s; keeping session alive: %s", - context, - session_key, - exc, + context, session_key, exc, ) return True - - - - - - - - - - - def has_any_sessions(self) -> bool: """Whether any session has ever been created (across all platforms). @@ -1116,12 +999,9 @@ class SessionStore( with inflight_lock: slot = self._inflight_sessions.get(session_key) - if slot is None: - slot = _SessionFlight() - self._inflight_sessions[session_key] = slot - owner = True - else: - owner = False + owner = slot is None + if owner: + slot = self._inflight_sessions[session_key] = _SessionFlight() if not owner: slot.event.wait() @@ -1133,13 +1013,10 @@ class SessionStore( return slot.result try: - result = self._get_or_create_session_impl( - source, - force_new=force_new, - touch_activity=touch_activity, + slot.result = self._get_or_create_session_impl( + source, force_new=force_new, touch_activity=touch_activity, ) - slot.result = result - return result + return slot.result except BaseException as exc: slot.error = exc raise @@ -1148,7 +1025,6 @@ class SessionStore( with inflight_lock: self._inflight_sessions.pop(session_key, None) - def _get_or_create_session_impl( self, source: SessionSource, @@ -1166,159 +1042,167 @@ class SessionStore( if not force_new: self._adopt_legacy_slack_entry(source, session_key) - db_end_session_id = None - db_create_kwargs = None - force_new_observed_entry = None - - # ---- Phase 1: lock read -- entry snapshot for stale/reset checks ---- - _stale_session_id = None - _entry_for_checks = None + # Phase 1 (lock): snapshot the entry for stale/reset checks. with self._lock: self._ensure_loaded_locked() - if force_new: - force_new_observed_entry = self._entries.get(session_key) - elif session_key in self._entries: - _entry_for_checks = self._entries[session_key] - _stale_session_id = _entry_for_checks.session_id + observed = self._entries.get(session_key) + # Phase 1b (no lock): compression tip + stale check + reset policy. + checks = None if force_new or observed is None else self._route_checks(observed, source, now) + # Phase 2 (lock): apply the decisions to _entries. + decision = self._apply_route_checks(session_key, checks, force_new, touch_activity, now) - # ---- Phase 1b: no-lock I/O -- compression tip + stale check + reset policy ---- - canonical_existing_session_id = None - _is_stale = False - _reset_reason = None - if _entry_for_checks is not None: - canonical_existing_session_id = self._compression_tip_for_session_id(_stale_session_id) - _is_stale = self._is_session_ended_in_db(_stale_session_id) - _reset_reason = self._route_reset_reason(_entry_for_checks, source, now) + # Phase 3 (no lock): recovery + create + save + DB ops. + if decision.needs_recover and decision.prev_session_id is None: + self._route_recover(decision, session_key, source, now) + create_kwargs = None + if decision.entry is None: + create_kwargs = self._route_create(decision, session_key, source, now, force_new, observed) - # ---- Phase 2: lock write -- apply decisions to _entries ---- - _needs_save = False - # Healthy-path saves take the single-row UPSERT fast path; structural - # transitions (recover/create) keep the full rewrite. - _metadata_only_save = False - _needs_recover = False - entry: Optional[SessionEntry] = None - was_auto_reset = False - auto_reset_reason = None - reset_had_activity = False - prev_session_id: Optional[str] = None - - with self._lock: - self._ensure_loaded_locked() - - if session_key in self._entries and not force_new: - entry = self._entries[session_key] - # A heal rewrites entry.session_id, so it must reach the - # sessions.json mirror too (forces the full-rewrite save). - _healed = self._heal_compression_tip_locked( - entry, _stale_session_id, canonical_existing_session_id - ) - # If another thread replaced the entry during our lock-free - # window, the stale/reset decisions no longer apply: healthy. - _checked = entry.session_id == _stale_session_id - _stale_hit = _is_stale and _checked - if _stale_hit: - # Stale routing self-heal: the entry points at a session - # ALREADY ended in state.db. Drop it and fall through to - # recovery (reopens agent_close / ws_orphan_reap rows, - # fresh session for other end_reasons). - logger.warning( - "gateway.session: routing key %r -> %s is ended in " - "state.db but still live in sessions.json; dropping " - "stale entry and recovering/recreating the session " - "(#54878)", - session_key, entry.session_id, - ) - if _stale_hit or (_checked and _reset_reason): - # Honour an expiry/reset decision instead of silently - # reopening the session via recovery. - if _reset_reason: - was_auto_reset = True - auto_reset_reason = _reset_reason - reset_had_activity = entry.last_prompt_tokens > 0 - db_end_session_id = entry.session_id - prev_session_id = entry.session_id - self._entries.pop(session_key, None) - entry = None - _needs_recover = True - else: - # Internal/system events preserve the user-activity clock. - if touch_activity: - entry.updated_at = now - _needs_save = touch_activity or _healed - _metadata_only_save = touch_activity and not _healed - elif not force_new: - _needs_recover = True - - # ---- Phase 3: no-lock I/O -- recovery + create + save + DB ops ---- - if _needs_recover and db_end_session_id is None: - recovered = self._query_recoverable_session( - session_key=session_key, source=source, now=now, - ) - if recovered is not None: - recovered_reset_reason = self._should_reset(recovered, source) - if recovered_reset_reason: - was_auto_reset = True - auto_reset_reason = recovered_reset_reason - reset_had_activity = recovered.reset_had_activity - db_end_session_id = recovered.session_id - prev_session_id = recovered.session_id - else: - self._reopen_session_row(session_key, recovered.session_id) - with self._lock: - entry = self._entries.setdefault(session_key, recovered) - _needs_save = True - - if entry is None: - # Create a candidate outside the lock, then publish only if another - # worker has not already populated this routing key. - session_id = _new_session_id(now) - candidate = SessionEntry( - session_key=session_key, - session_id=session_id, - created_at=now, - updated_at=now, - origin=source, - display_name=source.chat_name, - platform=source.platform, - chat_type=source.chat_type, - was_auto_reset=was_auto_reset, - auto_reset_reason=auto_reset_reason, - reset_had_activity=reset_had_activity, - prev_session_id=prev_session_id, - ) - with self._lock: - current = self._entries.get(session_key) - if current is None or (force_new and current is force_new_observed_entry): - self._entries[session_key] = candidate - current = candidate - entry = current - _needs_save = True - if entry is candidate: - db_create_kwargs = self._session_create_kwargs( - session_id=session_id, - session_key=session_key, - origin=source, - source_value=source.platform.value, - display_name=source.chat_name, - parent_session_id=prev_session_id, - ) - - if _needs_save: - if _metadata_only_save: + if decision.needs_save: + if decision.metadata_only_save: self._save_entry(session_key) else: self._save_entries() self._finish_route_transition( session_key, - end_session_id=db_end_session_id, - end_reason=auto_reset_reason if auto_reset_reason else "session_reset", - create_kwargs=db_create_kwargs, + end_session_id=decision.prev_session_id, + end_reason=decision.reset_reason or "session_reset", + create_kwargs=create_kwargs, origin=source, - display_name=entry.display_name, + display_name=decision.entry.display_name, ) - return entry + return decision.entry + def _route_checks(self, entry: SessionEntry, source: SessionSource, now: datetime) -> _RouteChecks: + """Lock-free DB/config I/O for an existing route.""" + sid = entry.session_id + canonical = self._compression_tip_for_session_id(sid) + is_stale = self._is_session_ended_in_db(sid) + return _RouteChecks(sid, canonical, is_stale, self._route_reset_reason(entry, source, now)) + + def _apply_route_checks( + self, + session_key: str, + checks: Optional[_RouteChecks], + force_new: bool, + touch_activity: bool, + now: datetime, + ) -> _RouteDecision: + """Apply stale/reset decisions to ``_entries`` under ``_lock``. + + If another thread replaced the entry during the lock-free window, the + snapshot's decisions no longer apply and the route is treated as healthy. + """ + decision = _RouteDecision() + with self._lock: + self._ensure_loaded_locked() + if force_new: + return decision + entry = self._entries.get(session_key) + if entry is None: + decision.needs_recover = True + return decision + snapshot_sid = checks.session_id if checks else None + # A heal rewrites entry.session_id, so it must reach the + # sessions.json mirror too (forces the full-rewrite save). + healed = self._heal_compression_tip_locked( + entry, snapshot_sid, checks.canonical_id if checks else None + ) + checked = entry.session_id == snapshot_sid + stale_hit = checked and checks.is_stale + reset_reason = checks.reset_reason if checked else None + if stale_hit: + # Stale routing self-heal: the entry points at a session + # ALREADY ended in state.db. Drop it and fall through to + # recovery (reopens agent_close / ws_orphan_reap rows, + # fresh session for other end_reasons). + logger.warning( + "gateway.session: routing key %r -> %s is ended in " + "state.db but still live in sessions.json; dropping " + "stale entry and recovering/recreating the session " + "(#54878)", + session_key, entry.session_id, + ) + if stale_hit or reset_reason: + # Honour an expiry/reset decision instead of silently + # reopening the session via recovery. + if reset_reason: + decision.reset_reason = reset_reason + decision.reset_had_activity = entry.last_prompt_tokens > 0 + decision.prev_session_id = entry.session_id + self._entries.pop(session_key, None) + decision.needs_recover = True + else: + # Internal/system events preserve the user-activity clock. + if touch_activity: + entry.updated_at = now + decision.entry = entry + decision.needs_save = touch_activity or healed + decision.metadata_only_save = touch_activity and not healed + return decision + + def _route_recover( + self, decision: _RouteDecision, session_key: str, source: SessionSource, now: datetime + ) -> None: + """Adopt a recoverable state.db row, or schedule its reset (no lock held on entry).""" + recovered = self._query_recoverable_session(session_key=session_key, source=source, now=now) + if recovered is None: + return + reset_reason = self._should_reset(recovered, source) + if reset_reason: + decision.reset_reason = reset_reason + decision.reset_had_activity = recovered.reset_had_activity + decision.prev_session_id = recovered.session_id + return + self._reopen_session_row(session_key, recovered.session_id) + with self._lock: + decision.entry = self._entries.setdefault(session_key, recovered) + decision.needs_save = True + + def _route_create( + self, + decision: _RouteDecision, + session_key: str, + source: SessionSource, + now: datetime, + force_new: bool, + observed: Optional[SessionEntry], + ) -> Optional[Dict[str, Any]]: + """Create a candidate outside the lock, publish it only if another worker + has not already populated this routing key; returns ``create_session`` + kwargs when the candidate won.""" + session_id = _new_session_id(now) + candidate = SessionEntry( + session_key=session_key, + session_id=session_id, + created_at=now, + updated_at=now, + origin=source, + display_name=source.chat_name, + platform=source.platform, + chat_type=source.chat_type, + was_auto_reset=decision.reset_reason is not None, + auto_reset_reason=decision.reset_reason, + reset_had_activity=decision.reset_had_activity, + prev_session_id=decision.prev_session_id, + ) + with self._lock: + current = self._entries.get(session_key) + if current is None or (force_new and current is observed): + self._entries[session_key] = current = candidate + decision.entry = current + decision.needs_save = True + if current is not candidate: + return None + return self._session_create_kwargs( + session_id=session_id, + session_key=session_key, + origin=source, + source_value=source.platform.value, + display_name=source.chat_name, + parent_session_id=decision.prev_session_id, + ) def update_session( self, @@ -1341,19 +1225,15 @@ class SessionStore( entry.last_prompt_tokens = last_prompt_tokens # Snapshot peer fields under _lock so a concurrent reset/heal # cannot produce a torn peer row. - peer_session_id = entry.session_id - peer_origin = entry.origin - peer_display_name = entry.display_name + peer_session_id, peer_origin, peer_display_name = ( + entry.session_id, entry.origin, entry.display_name + ) # Metadata-only: single-row UPSERT, outside ``_lock``. self._save_entry(session_key) self._record_gateway_session_peer( - peer_session_id, - session_key, - peer_origin, - display_name=peer_display_name, + peer_session_id, session_key, peer_origin, display_name=peer_display_name, ) - def get_session_metadata(self, session_key: str, key: str, default: Any = None) -> Any: """Return a metadata value stored on a live session entry.""" with self._lock: @@ -1387,7 +1267,6 @@ class SessionStore( entry = self._entry_locked(session_key) return dict(entry.model_override) if entry and entry.model_override else None - def reset_session(self, session_key: str, display_name: Optional[str] = None) -> Optional[SessionEntry]: """Force reset a session, creating a new session ID.""" with self._lock: @@ -1436,41 +1315,34 @@ class SessionStore( self._save() return new_entry - def switch_session(self, session_key: str, target_session_id: str) -> Optional[SessionEntry]: """Point a session key at an existing session ID (``/resume``): ends the current row, reopens the target so resume matches the CLI.""" - db_end_session_id = None - new_entry = None - with self._lock: old_entry = self._entry_locked(session_key) if old_entry is None: return None if old_entry.session_id == target_session_id: return old_entry - db_end_session_id = old_entry.session_id new_entry = self._replace_route_locked( session_key, old_entry, target_session_id, _now(), display_name=old_entry.display_name, ) - if self._db_for_key(session_key) and db_end_session_id: + if self._db_for_key(session_key) and old_entry.session_id: self._promote_session_reset( - session_key, db_end_session_id, "session_switch", + session_key, old_entry.session_id, "session_switch", log=lambda e: logger.debug("Session DB end_session failed: %s", e), ) - if self._db_for_key(session_key): self._reopen_session_row(session_key, target_session_id, log_prefix="Session DB reopen_session failed") self._record_gateway_session_peer( target_session_id, session_key, - new_entry.origin if new_entry else None, - display_name=new_entry.display_name if new_entry else None, + new_entry.origin, + display_name=new_entry.display_name, include_compression_ancestors=True, ) - return new_entry def list_sessions(self, active_minutes: Optional[int] = None) -> List[SessionEntry]: @@ -1508,13 +1380,6 @@ class SessionStore( return entry.session_id if entry else None - - # Max in-memory pending messages per session (DB persistently broken). - - - - - def build_session_context( source: SessionSource, config: GatewayConfig, @@ -1522,11 +1387,7 @@ def build_session_context( ) -> SessionContext: """Build a full session context (for system prompt injection).""" connected = config.get_connected_platforms() - home_channels = {} - for platform in connected: - home = config.get_home_channel(platform) - if home: - home_channels[platform] = home + home_channels = {p: home for p in connected if (home := config.get_home_channel(p))} context = SessionContext( source=source, connected_platforms=connected, diff --git a/gateway/session_context.py b/gateway/session_context.py index a09fabdf00..8ba6df79e0 100644 --- a/gateway/session_context.py +++ b/gateway/session_context.py @@ -1,32 +1,23 @@ """Session-scoped context variables for the Hermes gateway. -Replaces the old ``os.environ``-based session state (``HERMES_SESSION_*``) -with ``contextvars.ContextVar``. The gateway processes messages concurrently -via asyncio; ``os.environ`` is process-global, so message B silently -overwrote message A's thread id before A's agent finished and notifications -routed to the wrong thread. ContextVar values are task-local (inherited by -``run_in_executor`` threads), so concurrent messages never interfere. - -``get_session_env(name, default="")`` is a drop-in for -``os.getenv("HERMES_SESSION_*", default)`` at existing tool call sites. +Replaces the old ``os.environ``-based ``HERMES_SESSION_*`` state with task-local +``ContextVar``s (inherited by ``run_in_executor`` threads), so concurrently handled +messages no longer clobber each other's routing ids. ``get_session_env`` is a drop-in +for ``os.getenv``. """ +import os from contextlib import contextmanager from contextvars import ContextVar from typing import Any, Iterator -# Distinguishes "never set in this context" (fall back to os.environ for -# CLI/cron compat) from "explicitly set to empty" by clear_session_vars (no fallback). +# "Never set in this context" (fall back to os.environ for CLI/cron compat), as distinct +# from "explicitly set to empty" by clear_session_vars (no fallback). _UNSET: Any = object() -# Process-level monotonic latch: has any code in this process bound a session via -# set_session_vars()? Concurrent multi-session hosts (gateway, ACP, API server, -# TUI, cron) do; a pure single-process CLI/one-shot does not. The subprocess-env -# bridge (tools/environments/local.py) reads this to pick its leak policy: when -# engaged, the ContextVars are authoritative and an _UNSET var means "no session -# bound in THIS task", so the last-writer-wins os.environ mirror must NOT be -# inherited by a child process. When never engaged, the os.environ fallback is -# preserved (no concurrency to leak across). +# Process-level monotonic latch: has any code bound a session via set_session_vars()? The +# subprocess-env bridge reads it: when engaged, ContextVars are authoritative and an _UNSET +# var means "no session bound in THIS task", so the os.environ mirror must NOT leak to a child. _session_context_engaged: bool = False @@ -35,121 +26,69 @@ def session_context_engaged() -> bool: return _session_context_engaged -def _var(name: str) -> ContextVar: - return ContextVar(name, default=_UNSET) - - # --- Per-task session variables -------------------------------------------- -_SESSION_PLATFORM = _var("HERMES_SESSION_PLATFORM") -_SESSION_SOURCE = _var("HERMES_SESSION_SOURCE") -_SESSION_CHAT_ID = _var("HERMES_SESSION_CHAT_ID") -_SESSION_CHAT_TYPE = _var("HERMES_SESSION_CHAT_TYPE") -_SESSION_CHAT_NAME = _var("HERMES_SESSION_CHAT_NAME") -_SESSION_THREAD_ID = _var("HERMES_SESSION_THREAD_ID") -_SESSION_USER_ID = _var("HERMES_SESSION_USER_ID") -_SESSION_USER_ID_ALT = _var("HERMES_SESSION_USER_ID_ALT") -_SESSION_USER_NAME = _var("HERMES_SESSION_USER_NAME") -# Platform-neutral scope discriminator (Discord guild / Slack workspace / Matrix -# server). Captured at bind time so async producers (delegate_task -# background=True, terminal watchers) can persist a completion's full routing -# origin: a relay connector's fail-closed egress guard needs scope_id (or a user -# binding) to resolve the tenant for a scoped reply after a restart. -_SESSION_SCOPE_ID = _var("HERMES_SESSION_SCOPE_ID") -_SESSION_KEY = _var("HERMES_SESSION_KEY") -_SESSION_ID = _var("HERMES_SESSION_ID") -# In-process UI tab/window id for multi-session desktop/TUI hosts — deliberately -# separate from the durable HERMES_SESSION_ID. Background completions use it as -# a precise return address so a stale/rotated durable key cannot be consumed by -# whichever desktop poller wakes first. -_SESSION_UI_SESSION_ID = _var("HERMES_UI_SESSION_ID") -# Triggering message id: reply anchor so background notifications stay inside -# the originating Telegram private-chat topic (routes only with thread id + anchor). -_SESSION_MESSAGE_ID = _var("HERMES_SESSION_MESSAGE_ID") -_SESSION_PROFILE = _var("HERMES_SESSION_PROFILE") -_BROWSER_CONTROL_PRINCIPAL = _var("HERMES_BROWSER_CONTROL_PRINCIPAL") -_BROWSER_CONTROL_TRANSPORT_FAMILY = _var("HERMES_BROWSER_CONTROL_TRANSPORT_FAMILY") -# Per-session cron marker, tri-state: _UNSET keeps the legacy env fallback for -# CLI/tests; "1" marks cron; "" explicitly marks non-cron and masks leaked env. -_CRON_SESSION = _var("HERMES_CRON_SESSION") - -# Whether this session's channel can route an ASYNC completion back to the agent -# AFTER the turn ends (wake a fresh turn). True for long-lived CLI sessions and -# real gateway platforms (persistent outbound channel + watcher/drain loops); -# False for finite runtimes that may exit before a detached completion returns -# (stateless API-server requests, dispatcher-spawned Kanban workers). Tools that -# promise async delivery (terminal notify_on_complete / watch_patterns, -# delegate_task background=True) read ``async_delivery_supported()`` and refuse a -# promise the channel can't keep. Default _UNSET => supported, so the CLI (never -# sets a platform) and contextvar-unaware paths keep working; stateless adapters -# opt OUT via ``supports_async_delivery = False`` on the adapter class, which the -# gateway propagates here at session-bind time. -_SESSION_ASYNC_DELIVERY = _var("HERMES_SESSION_ASYNC_DELIVERY") - -# Cron auto-delivery vars — set per-job in run_job() so concurrent jobs don't -# clobber each other's delivery targets. -_CRON_AUTO_DELIVER_PLATFORM = _var("HERMES_CRON_AUTO_DELIVER_PLATFORM") -_CRON_AUTO_DELIVER_CHAT_ID = _var("HERMES_CRON_AUTO_DELIVER_CHAT_ID") -_CRON_AUTO_DELIVER_THREAD_ID = _var("HERMES_CRON_AUTO_DELIVER_THREAD_ID") - -# Vars bound by set_session_vars / cleared to "" by clear_session_vars, in order. +# Bound by set_session_vars / cleared to "" by clear_session_vars; tuple ORDER is the +# positional order of ``values`` in set_session_vars (they are zipped). +# * SCOPE_ID: platform-neutral scope (guild / workspace / Matrix server), captured so async +# producers can persist a completion's full routing origin (relay egress guards need it). +# * UI_SESSION_ID: in-process UI tab id, separate from the durable SESSION_ID: a precise +# return address so a stale/rotated durable key is not consumed by the wrong poller. +# * MESSAGE_ID: reply anchor keeping background notifications inside the originating +# Telegram private-chat topic. +# * CRON_SESSION: tri-state — _UNSET keeps the legacy env fallback for CLI/tests; "1" +# marks cron; "" explicitly marks non-cron and masks leaked env. _SESSION_VARS = ( - _SESSION_PLATFORM, - _SESSION_SOURCE, - _SESSION_CHAT_ID, - _SESSION_CHAT_TYPE, - _SESSION_CHAT_NAME, - _SESSION_THREAD_ID, - _SESSION_USER_ID, - _SESSION_USER_ID_ALT, - _SESSION_USER_NAME, - _SESSION_SCOPE_ID, - _SESSION_KEY, - _SESSION_ID, - _SESSION_UI_SESSION_ID, - _SESSION_MESSAGE_ID, - _SESSION_PROFILE, - _BROWSER_CONTROL_PRINCIPAL, - _BROWSER_CONTROL_TRANSPORT_FAMILY, - _CRON_SESSION, -) + _SESSION_PLATFORM, _SESSION_SOURCE, _SESSION_CHAT_ID, _SESSION_CHAT_TYPE, + _SESSION_CHAT_NAME, _SESSION_THREAD_ID, _SESSION_USER_ID, _SESSION_USER_ID_ALT, + _SESSION_USER_NAME, _SESSION_SCOPE_ID, _SESSION_KEY, _SESSION_ID, + _SESSION_UI_SESSION_ID, _SESSION_MESSAGE_ID, _SESSION_PROFILE, + _BROWSER_CONTROL_PRINCIPAL, _BROWSER_CONTROL_TRANSPORT_FAMILY, _CRON_SESSION, +) = tuple(ContextVar(name, default=_UNSET) for name in ( + "HERMES_SESSION_PLATFORM", "HERMES_SESSION_SOURCE", "HERMES_SESSION_CHAT_ID", + "HERMES_SESSION_CHAT_TYPE", "HERMES_SESSION_CHAT_NAME", "HERMES_SESSION_THREAD_ID", + "HERMES_SESSION_USER_ID", "HERMES_SESSION_USER_ID_ALT", "HERMES_SESSION_USER_NAME", + "HERMES_SESSION_SCOPE_ID", "HERMES_SESSION_KEY", "HERMES_SESSION_ID", + "HERMES_UI_SESSION_ID", "HERMES_SESSION_MESSAGE_ID", "HERMES_SESSION_PROFILE", + "HERMES_BROWSER_CONTROL_PRINCIPAL", "HERMES_BROWSER_CONTROL_TRANSPORT_FAMILY", + "HERMES_CRON_SESSION", +)) -# Legacy env-var name -> ContextVar, for get_session_env. _SESSION_ASYNC_DELIVERY -# is deliberately absent: it is a bool capability read via async_delivery_supported. -_VAR_MAP = { - var.name: var - for var in (*_SESSION_VARS, _CRON_AUTO_DELIVER_PLATFORM, _CRON_AUTO_DELIVER_CHAT_ID, _CRON_AUTO_DELIVER_THREAD_ID) -} +# Whether this channel can route an ASYNC completion back AFTER the turn ends (read via +# ``async_delivery_supported()``). False for finite runtimes that may exit first (stateless +# API-server requests, Kanban workers). Default _UNSET => supported, so the CLI and +# contextvar-unaware paths keep working; stateless adapters opt OUT via +# ``supports_async_delivery = False``, propagated at bind time. +_SESSION_ASYNC_DELIVERY = ContextVar("HERMES_SESSION_ASYNC_DELIVERY", default=_UNSET) + +# Cron auto-delivery vars, set per-job in run_job() so concurrent jobs don't clobber. +_CRON_AUTO_DELIVER_PLATFORM = ContextVar("HERMES_CRON_AUTO_DELIVER_PLATFORM", default=_UNSET) +_CRON_AUTO_DELIVER_CHAT_ID = ContextVar("HERMES_CRON_AUTO_DELIVER_CHAT_ID", default=_UNSET) +_CRON_AUTO_DELIVER_THREAD_ID = ContextVar("HERMES_CRON_AUTO_DELIVER_THREAD_ID", default=_UNSET) + +# Legacy env-var name -> ContextVar, for get_session_env. _SESSION_ASYNC_DELIVERY is +# deliberately absent: it is a bool capability read via async_delivery_supported. +_VAR_MAP = {var.name: var for var in ( + *_SESSION_VARS, _CRON_AUTO_DELIVER_PLATFORM, _CRON_AUTO_DELIVER_CHAT_ID, + _CRON_AUTO_DELIVER_THREAD_ID, +)} -def _clear_session_cwd() -> None: +def _runtime_cwd(func: str, *args: Any) -> None: + """Best-effort call of ``agent.runtime_cwd.``; import/runtime failures are ignored.""" try: - from agent.runtime_cwd import clear_session_cwd - - clear_session_cwd() + from agent import runtime_cwd + getattr(runtime_cwd, func)(*args) except Exception: pass def set_current_session_id(session_id: str) -> None: - """Synchronize ``HERMES_SESSION_ID`` across ContextVar and ``os.environ``. - - Long-lived single-process entrypoints (CLI) rotate sessions via /new, - /resume, /branch, or compression splits without rebuilding the agent; tools - read ``get_session_env("HERMES_SESSION_ID")`` with an os.environ fallback, - so both stores must move together. - - Delegated subagent children are the exception: they are built in the parent - process inside ``delegated_child_context()`` and their ``AIAgent.__init__`` - calls this helper. Writing the child's id to process-global ``os.environ`` - would clobber the parent's id for the rest of the process, so only the - task-local ContextVar write happens for them. Root agents keep both paths. - """ - import os - + """Synchronize ``HERMES_SESSION_ID`` across ContextVar and ``os.environ`` (tools read it + with an os.environ fallback). Delegated subagent children (built in the parent process) + get ONLY the task-local write, or they would clobber the parent's id.""" _SESSION_ID.set(session_id) try: from agent.delegation_context import is_delegated_child_context - if is_delegated_child_context(): return except Exception: @@ -159,12 +98,8 @@ def set_current_session_id(session_id: str) -> None: @contextmanager def scoped_current_session_id(session_id: str | None = None) -> Iterator[None]: - """Bind a task-local session id and restore the prior value on exit. - - ``session_id=None`` makes this a pure save/restore boundary around code that - may call :func:`set_current_session_id` itself (delegated ``AIAgent`` - construction). Never mutates ``os.environ``. - """ + """Bind a task-local session id and restore the prior value on exit; never touches + ``os.environ``. ``session_id=None`` is a pure save/restore boundary.""" previous = _SESSION_ID.get() if session_id is not None: _SESSION_ID.set(session_id) @@ -175,188 +110,90 @@ def scoped_current_session_id(session_id: str | None = None) -> Iterator[None]: def set_session_vars( - platform: str = "", - source: str = "", - chat_id: str = "", - chat_type: str = "", - chat_name: str = "", - thread_id: str = "", - user_id: str = "", - user_id_alt: str = "", - user_name: str = "", - scope_id: str = "", - session_key: str = "", - session_id: str = "", - message_id: str = "", - profile: str = "", - browser_control_principal: str = "", - browser_control_transport_family: str = "", - cwd: str = "", - async_delivery: bool = True, - ui_session_id: str = "", - cron_session: Any = _UNSET, + platform: str = "", source: str = "", chat_id: str = "", chat_type: str = "", + chat_name: str = "", thread_id: str = "", user_id: str = "", user_id_alt: str = "", + user_name: str = "", scope_id: str = "", session_key: str = "", session_id: str = "", + message_id: str = "", profile: str = "", browser_control_principal: str = "", + browser_control_transport_family: str = "", cwd: str = "", async_delivery: bool = True, + ui_session_id: str = "", cron_session: Any = _UNSET, ) -> list: - """Set all session context variables and return reset tokens. - - Call ``clear_session_vars(tokens)`` in a ``finally`` when the handler exits. - These helpers are not nestable: clearing resets every var to ``""`` rather - than restoring prior values, and the tokens are accepted only for API compat. - - ``cwd`` pins the logical working directory. ``async_delivery`` declares - whether the channel can route a background completion back after the turn - (stateless adapters such as the API server pass ``False``). ``cron_session`` - is tri-state; see ``_CRON_SESSION``. - """ - # Latch the process as engaged — see _session_context_engaged. + """Set all session context variables and return reset tokens. Call + ``clear_session_vars(tokens)`` in a ``finally``; not nestable, clearing resets every var + to ``""`` rather than restoring prior values (tokens are accepted only for API compat).""" global _session_context_engaged _session_context_engaged = True values = ( - platform, source, chat_id, chat_type, chat_name, thread_id, user_id, - user_id_alt, user_name, scope_id, session_key, session_id, ui_session_id, - message_id, profile, browser_control_principal, - browser_control_transport_family, cron_session, + platform, source, chat_id, chat_type, chat_name, thread_id, user_id, user_id_alt, + user_name, scope_id, session_key, session_id, ui_session_id, message_id, profile, + browser_control_principal, browser_control_transport_family, cron_session, ) tokens = [var.set(value) for var, value in zip(_SESSION_VARS, values)] tokens.append(_SESSION_ASYNC_DELIVERY.set(bool(async_delivery))) - try: - from agent.runtime_cwd import set_session_cwd - - set_session_cwd(cwd) - except Exception: - pass + _runtime_cwd("set_session_cwd", cwd) return tokens def clear_session_vars(tokens: list) -> None: - """Mark session context variables as explicitly cleared. - - Sets every var to ``""`` (not ``var.reset(token)``) so ``get_session_env`` - returns empty instead of falling back to stale ``os.environ`` values while - staying distinguishable from "never set" (``_UNSET``). Async-delivery is - reset to ``_UNSET`` rather than a falsy value: a cleared context must fall - back to default-supported, not look like an opted-out stateless adapter. - """ + """Mark session context variables as explicitly cleared (``""``, not ``_UNSET``), so + ``get_session_env`` returns empty instead of stale ``os.environ`` values. Async-delivery + goes back to ``_UNSET``: a cleared context is default-supported, not opted-out.""" for var in _SESSION_VARS: var.set("") _SESSION_ASYNC_DELIVERY.set(_UNSET) - _clear_session_cwd() + _runtime_cwd("clear_session_cwd") def reset_session_vars() -> None: - """Reset every session context variable to ``_UNSET`` for THIS context. - - Unlike :func:`clear_session_vars` (``""`` = "explicitly cleared", used when a - handler *finishes*), this restores "never bound here" — what a freshly - spawned task should look like *before* binding its own session. - - Why: ``create_task`` snapshots the current context, so message B's task can - inherit message A's already-**set** vars. Until B binds its own session, - any subprocess it spawns reads A's identity through the subprocess-env - bridge — whose _UNSET-strip guard cannot help because the vars are set-to-A. - Calling this at the top of the per-message handler makes that window strip - safe (no session) instead of leaking the foreign one. See - tests/tools/test_local_env_session_leak.py and - tests/gateway/test_session_context_inheritance.py. - - ``_SESSION_ASYNC_DELIVERY`` is reset explicitly (it lives outside - ``_VAR_MAP``): otherwise a task spawned from a context where a sibling - adapter bound ``async_delivery=False`` inherits that ``False`` through the - pre-bind window and misreports the new channel as unable to deliver. - """ + """Reset every session var to ``_UNSET`` ("never bound here") for THIS context. Call at + the top of a fresh task *before* it binds: ``create_task`` snapshots the context, so B's + task inherits A's already-set vars and a subprocess spawned before B binds would read A's + identity. ``_SESSION_ASYNC_DELIVERY`` (outside ``_VAR_MAP``) is reset explicitly too.""" for var in _VAR_MAP.values(): var.set(_UNSET) _SESSION_ASYNC_DELIVERY.set(_UNSET) - _clear_session_cwd() + _runtime_cwd("clear_session_cwd") def get_session_env(name: str, default: str = "") -> str: - """Read a session context variable by its legacy ``HERMES_SESSION_*`` name. - - Drop-in for ``os.getenv(name, default)``. Resolution: the ContextVar if it - was ever set in this context (even to ``""`` — no fallback); else - ``os.environ`` (CLI, cron scheduler, tests that never bind); else *default*. - """ - import os - + """Read a session var by legacy ``HERMES_SESSION_*`` name; drop-in for os.getenv. The + ContextVar wins if ever set here (even to ``""``); else ``os.environ``; else *default*.""" var = _VAR_MAP.get(name) if var is not None and (value := var.get()) is not _UNSET: return value return os.getenv(name, default) -# Surfaces that are not a human chat channel. The gateway binds a platform value -# (``telegram``) to HERMES_SESSION_PLATFORM while the CLI/TUI/desktop bind -# HERMES_SESSION_SOURCE and leave platform empty, so both are consulted. -# ``local``, ``api_server``, ``webhook``, ``msgraph_webhook`` are real Platform -# values with no attachment channel behind them. Default-deny: an unrecognized -# identity counts as messaging so a new chat platform is never treated as a -# private surface before this set is updated. Mirrors LOCAL_SESSION_SOURCE_IDS -# in apps/desktop/src/lib/session-source.ts; keep roughly in sync. -NON_MESSAGING_SESSION_SURFACES = frozenset( - { - "", - "api_server", - "cli", - "codex", - "desktop", - "gateway", - "kanban", - "local", - "msgraph_webhook", - "tool", - "tui", - "webhook", - } -) +# Surfaces that are not a human chat channel (the gateway binds HERMES_SESSION_PLATFORM, +# CLI/TUI/desktop bind HERMES_SESSION_SOURCE, so both are consulted). ``local``, +# ``api_server``, ``webhook``, ``msgraph_webhook`` are real Platform values with no +# attachment channel. Default-deny: an unrecognized identity counts as messaging. +# Mirrors LOCAL_SESSION_SOURCE_IDS in apps/desktop/src/lib/session-source.ts. +NON_MESSAGING_SESSION_SURFACES = frozenset({ + "", "api_server", "cli", "codex", "desktop", "gateway", "kanban", "local", + "msgraph_webhook", "tool", "tui", "webhook", +}) def session_is_messaging_surface() -> bool: - """Whether this turn is delivered over a human messaging channel. - - Decides "user is reading a chat message" vs "user is at a machine they own": - delivery tags, whether a file must land somewhere the gateway can send from, - whether narration reads as chat noise. Checks ``HERMES_PLATFORM``, then the - session platform, then the session source against - :data:`NON_MESSAGING_SESSION_SURFACES`. - """ - import os - + """Whether this turn is delivered over a human messaging channel (checks + ``HERMES_PLATFORM``, then the session platform, then the session source).""" platform = os.getenv("HERMES_PLATFORM") or get_session_env("HERMES_SESSION_PLATFORM", "") - source = get_session_env("HERMES_SESSION_SOURCE", "") - return any( - (ident := str(identity or "").strip().lower()) and ident not in NON_MESSAGING_SESSION_SURFACES - for identity in (platform, source) - ) + idents = (platform, get_session_env("HERMES_SESSION_SOURCE", "")) + idents = (str(v or "").strip().lower() for v in idents) + return any(ident and ident not in NON_MESSAGING_SESSION_SURFACES for ident in idents) def declare_stateless_channel() -> None: - """Declare that this session cannot receive an async background completion. - - Binds only the delivery capability. Use this instead of - ``set_session_vars(async_delivery=False)`` on a pure single-process runner: - ``set_session_vars`` also latches ``_session_context_engaged``, which flips - the subprocess env bridge to ContextVar-authoritative — a one-shot CLI must - not flip that latch as a side effect of declaring a capability. Callers that - build a full context (cron's ``run_job``) pass ``async_delivery=False``. - ``delegate_task`` then falls through to its inline path so results return - within the turn instead of going to a channel that never delivers. - """ + """Declare that this session cannot receive an async background completion. Unlike + ``set_session_vars(async_delivery=False)`` this does NOT latch ``_session_context_engaged`` + (flipping the subprocess env bridge), which a one-shot CLI must not do as a side effect.""" _SESSION_ASYNC_DELIVERY.set(False) def async_delivery_supported() -> bool: - """Whether the current session can deliver a background completion later. - - False for finite runtimes: sessions bound by a stateless channel (API - server, ``hermes -z``, cron — see :func:`declare_stateless_channel`) and - dispatcher-spawned Kanban workers (``HERMES_KANBAN_TASK``), which are - one-shot ``chat -q`` subprocesses whose parent disappears after the quiet - turn, so a later completion has no durable consumer. Gateway platforms, - the interactive CLI, and any path that never bound the var return True. - """ - import os - - # Kanban worker: force tools onto their synchronous/polling fallbacks. + """Whether the current session can deliver a background completion later. False for + stateless channels (:func:`declare_stateless_channel`) and Kanban workers + (``HERMES_KANBAN_TASK``: one-shot subprocesses whose parent disappears after the turn).""" if os.environ.get("HERMES_KANBAN_TASK"): return False value = _SESSION_ASYNC_DELIVERY.get() diff --git a/gateway/session_db_recovery.py b/gateway/session_db_recovery.py index 0315049969..2c00184223 100644 --- a/gateway/session_db_recovery.py +++ b/gateway/session_db_recovery.py @@ -2,14 +2,13 @@ from __future__ import annotations +import contextlib import threading import time import weakref from dataclasses import dataclass from pathlib import Path from typing import Any, Callable -import contextlib - _INITIAL_RETRY_DELAY_SECONDS = 1.0 _MAX_RETRY_DELAY_SECONDS = 60.0 @@ -27,9 +26,8 @@ class _HealthSource: _health_lock = threading.Lock() -_health_states: weakref.WeakKeyDictionary[_HealthSource, dict[Path, str]] = ( - weakref.WeakKeyDictionary() -) +_health_states: weakref.WeakKeyDictionary[_HealthSource, dict[Path, str]] +_health_states = weakref.WeakKeyDictionary() def _publish_health(source: _HealthSource, path: Path, state: str) -> None: @@ -40,7 +38,6 @@ def _publish_health(source: _HealthSource, path: Path, state: str) -> None: aggregate = next((s for s in ("retrying", "unavailable") if s in all_states), "ok") try: from gateway.status import write_runtime_status - write_runtime_status(session_store={"status": aggregate}) except Exception: pass # Runtime health is diagnostic only; persistence must not depend on it. @@ -49,16 +46,12 @@ def _publish_health(source: _HealthSource, path: Path, state: str) -> None: class RecoverableHandleCache: """Cache handles by path while allowing failed opens to heal in-process. - Opens run OUTSIDE ``lock`` (single-flight per path via ``in_flight``); a - ``close_all`` bumps ``_generation`` so any open that completes afterwards is - treated as stale and rejected rather than resurrecting a drained cache. + Opens run OUTSIDE ``lock`` (single-flight per path via ``in_flight``); ``close_all`` bumps + ``_generation`` so a later-completing open is rejected instead of resurrecting the cache. """ def __init__( - self, - *, - handles: dict[Path, Any] | None = None, - lock: threading.Lock | None = None, + self, *, handles: dict[Path, Any] | None = None, lock: threading.Lock | None = None, clock: Callable[[], float] = time.monotonic, initial_retry_delay: float = _INITIAL_RETRY_DELAY_SECONDS, max_retry_delay: float = _MAX_RETRY_DELAY_SECONDS, @@ -78,20 +71,13 @@ class RecoverableHandleCache: return generation != self._generation or self._unavailable.get(path) is not unavailable def get( - self, - path: Path, - opener: Callable[[], Any], - *, - raise_on_error: bool = False, + self, path: Path, opener: Callable[[], Any], *, raise_on_error: bool = False, on_recovered: Callable[[], None] | None = None, non_cacheable: Callable[[Exception], bool] | None = None, ) -> Any: - """Return a cached handle or make one bounded, single-flight open attempt. - - Returns None while a retry is in flight or backing off (callers fall back). - ``non_cacheable`` exceptions (e.g. a live-system guard) are re-raised - without recording a failure so the next call retries immediately. - """ + """Return a cached handle or make one bounded, single-flight open attempt; None while + a retry is in flight or backing off. ``non_cacheable`` exceptions (e.g. a live-system + guard) are re-raised without recording a failure so the next call retries at once.""" path = Path(path) with self.lock: if path in self.handles: @@ -102,7 +88,6 @@ class RecoverableHandleCache: unavailable.in_flight = True was_unavailable = unavailable.failures > 0 generation = self._generation - if was_unavailable: _publish_health(self._health_source, path, "retrying") @@ -118,11 +103,8 @@ class RecoverableHandleCache: raise if not stale: unavailable.failures += 1 - delay = min( - self._initial_retry_delay * (2 ** min(unavailable.failures - 1, 30)), - self._max_retry_delay, - ) - unavailable.next_retry_at = self._clock() + delay + backoff = self._initial_retry_delay * (2 ** min(unavailable.failures - 1, 30)) + unavailable.next_retry_at = self._clock() + min(backoff, self._max_retry_delay) unavailable.in_flight = False if not stale: _publish_health(self._health_source, path, "unavailable") @@ -159,7 +141,6 @@ class RecoverableHandleCache: with contextlib.suppress(Exception): close(handle) with _health_lock: - states = _health_states.get(self._health_source) - if states is not None: - for path in paths: - states.pop(path, None) + states = _health_states.get(self._health_source, {}) + for path in paths: + states.pop(path, None) diff --git a/gateway/session_lifecycle.py b/gateway/session_lifecycle.py index 4497cce47d..287b5126a5 100644 --- a/gateway/session_lifecycle.py +++ b/gateway/session_lifecycle.py @@ -1,6 +1,6 @@ """SessionStore reset/expiry policy and crash-recovery markers: idle/daily reset evaluation, expiry finalization, active-turn tokens, resume_pending, -suspension and pruning. +suspension and pruning. Also home of the shared clock/id helpers. Mixin split out of ``gateway/session.py``; bound onto ``SessionStore`` via the MRO. """ @@ -8,6 +8,7 @@ Mixin split out of ``gateway/session.py``; bound onto ``SessionStore`` via the M from __future__ import annotations import logging +import os import uuid from datetime import datetime, timedelta from typing import TYPE_CHECKING, Optional @@ -19,11 +20,52 @@ if TYPE_CHECKING: logger = logging.getLogger("gateway.session") -class SessionLifecycleMixin: - """SessionStore reset/expiry policy and crash-recovery markers: idle/daily - reset evaluation, expiry finalization, active-turn tokens, resume_pending, - suspension and pruning. +def _now() -> datetime: + """Return the current local time.""" + return datetime.now() + + +def _new_session_id(now: datetime) -> str: + return f"{now.strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:8]}" + + +def _iso(dt: Optional[datetime]) -> Optional[str]: + return dt.isoformat() if dt else None + + +def _parse_iso(value) -> Optional[datetime]: + """``datetime.fromisoformat`` that returns None for empty/malformed input.""" + if not value: + return None + try: + return datetime.fromisoformat(value) + except (TypeError, ValueError): + return None + + +# Default auto-continue freshness window (1 hour): a restart-interrupted +# session is only auto-resumed while within this window of when +# ``resume_pending`` was marked. ``gateway/run.py`` bridges config.yaml +# ``agent.gateway_auto_continue_freshness`` into the env var at startup. +_AUTO_CONTINUE_FRESHNESS_SECS_DEFAULT = 60 * 60 + + +def auto_continue_freshness_window() -> float: + """Auto-continue freshness window in seconds (single source of truth for + the resume scheduler and the routing-time zombie gate). + + Reads ``HERMES_AUTO_CONTINUE_FRESHNESS``; falls back to the default when + unset or malformed. Non-positive disables the gate. """ + raw = os.environ.get("HERMES_AUTO_CONTINUE_FRESHNESS") + try: + return float(raw) if raw else float(_AUTO_CONTINUE_FRESHNESS_SECS_DEFAULT) + except (TypeError, ValueError): + return float(_AUTO_CONTINUE_FRESHNESS_SECS_DEFAULT) + + +class SessionLifecycleMixin: + """SessionStore reset/expiry policy and crash-recovery markers.""" def set_expiry_finalized( self, entry: SessionEntry, *, clear_model_override: bool = True @@ -43,35 +85,33 @@ class SessionLifecycleMixin: # Background caller never entered ``_profile_runtime_scope``: resolve # the store from the key, not the ambient scope. _db = self._db_for_key(entry.session_key) - if _db: - setter = getattr(_db, "set_expiry_finalized", None) - if callable(setter): - try: - setter(entry.session_id, True) - except Exception as exc: - logger.debug("Session DB expiry_finalized write failed for %s: %s", entry.session_id, exc) + if not _db: + return + setter = getattr(_db, "set_expiry_finalized", None) + if callable(setter): try: - # Without a durable ``session_reset`` end_reason, later agent - # cleanup ends the row as ``agent_close``, which stale-route - # recovery treats as resumable. Promotion only upgrades live/ - # agent_close rows; explicit boundaries are preserved. - _db.promote_to_session_reset(entry.session_id) + setter(entry.session_id, True) except Exception as exc: - logger.debug("Session DB promote_to_session_reset failed for %s: %s", entry.session_id, exc) + logger.debug("Session DB expiry_finalized write failed for %s: %s", entry.session_id, exc) + try: + # Without a durable ``session_reset`` end_reason, later agent + # cleanup ends the row as ``agent_close``, which stale-route + # recovery treats as resumable. Promotion only upgrades live/ + # agent_close rows; explicit boundaries are preserved. + _db.promote_to_session_reset(entry.session_id) + except Exception as exc: + logger.debug("Session DB promote_to_session_reset failed for %s: %s", entry.session_id, exc) @staticmethod def _policy_reset_reason(policy, updated_at: datetime) -> Optional[str]: """Return "idle"/"daily" when *updated_at* is overdue under *policy*, else None.""" - from gateway.session import _now if policy.mode == "none": return None now = _now() if policy.mode in {"idle", "both"} and now > updated_at + timedelta(minutes=policy.idle_minutes): return "idle" if policy.mode in {"daily", "both"}: - today_reset = now.replace( - hour=policy.at_hour, minute=0, second=0, microsecond=0, - ) + today_reset = now.replace(hour=policy.at_hour, minute=0, second=0, microsecond=0) if now.hour < policy.at_hour: today_reset -= timedelta(days=1) if updated_at < today_reset: @@ -87,10 +127,7 @@ class SessionLifecycleMixin: if self._has_active_processes_safe(entry.session_key, context="expiry"): logger.debug("Session %s not expired — active background processes", entry.session_key) return False - policy = self.config.get_reset_policy( - platform=entry.platform, - session_type=entry.chat_type, - ) + policy = self.config.get_reset_policy(platform=entry.platform, session_type=entry.chat_type) return self._policy_reset_reason(policy, entry.updated_at) is not None def is_session_finalizable(self, entry: SessionEntry) -> bool: @@ -102,10 +139,7 @@ class SessionLifecycleMixin: resolution errors count as "not finalizable" (sweep reaps — safe). """ try: - policy = self.config.get_reset_policy( - platform=entry.platform, - session_type=entry.chat_type, - ) + policy = self.config.get_reset_policy(platform=entry.platform, session_type=entry.chat_type) return policy.mode != "none" except Exception: return False @@ -114,8 +148,8 @@ class SessionLifecycleMixin: """True iff state.db has this session with a non-null end_reason. Same staleness test as ``_prune_stale_sessions_locked`` (no DB, no - row, or DB error -> False, keep). Used by ``get_or_create_session`` - to self-heal at routing time, since the startup prune cannot see a + row, or DB error -> False, keep). Lets ``get_or_create_session`` + self-heal at routing time, since the startup prune cannot see a session ended while the gateway stays alive. Store resolved from the row's owning profile, not the ambient scope. """ @@ -129,18 +163,13 @@ class SessionLifecycleMixin: return bool(row is not None and row.get("end_reason") is not None) def _should_reset(self, entry: SessionEntry, source: SessionSource) -> Optional[str]: - """Return the reset reason ("idle"/"daily") if policy says reset, else None. - - Sessions with active background processes are never reset. - """ + """Reset reason ("idle"/"daily") if policy says reset, else None. + Sessions with active background processes are never reset.""" session_key = self._generate_session_key(source) if self._has_active_processes_safe(session_key, context="reset"): logger.debug("Session reset skipped for %s — active background processes", session_key) return None - policy = self.config.get_reset_policy( - platform=source.platform, - session_type=source.chat_type - ) + policy = self.config.get_reset_policy(platform=source.platform, session_type=source.chat_type) return self._policy_reset_reason(policy, entry.updated_at) def _route_reset_reason( @@ -154,15 +183,12 @@ class SessionLifecycleMixin: makes an expired marker fall through to a normal resume, never a silent fresh session. """ - from gateway.session import auto_continue_freshness_window if entry.suspended: return "suspended" reason = self._should_reset(entry, source) if reason or not entry.resume_pending: return reason - policy = self.config.get_reset_policy( - platform=source.platform, session_type=source.chat_type, - ) + policy = self.config.get_reset_policy(platform=source.platform, session_type=source.chat_type) if policy.mode == "none": return None window = auto_continue_freshness_window() @@ -181,57 +207,60 @@ class SessionLifecycleMixin: self._save() return True + def _update_all_entries_locked(self, mutate) -> int: + """Apply ``mutate(entry) -> bool`` to every entry under ``_lock``; save once + if any returned True. Returns the count that did.""" + with self._lock: + self._ensure_loaded_locked() + changed = sum(1 for entry in self._entries.values() if mutate(entry)) + if changed: + self._save() + return changed + def suspend_session(self, session_key: str) -> bool: """Mark a session suspended so it auto-resets on next access (/stop). Returns True if the session existed.""" return self._update_entry(session_key, lambda e: setattr(e, "suspended", True)) + def _set_turn_marker_locked(self, session_key: str, entry: SessionEntry, token, started_at) -> None: + """Persist the active-turn pair BEFORE publishing it in memory, so a failed + write can neither leak an unowned token nor drop a live one. Lock held.""" + candidate = entry.to_dict() + candidate["active_turn_token"] = token + candidate["active_turn_started_at"] = _iso(started_at) + if started_at is not None: + # Keeps the legacy 120s startup heuristic working for an older + # binary during a rolling downgrade/upgrade window. + candidate["updated_at"] = started_at.isoformat() + self._save_entry(session_key, entry_data=candidate, lock_held=True) + entry.active_turn_token = token + entry.active_turn_started_at = started_at + if started_at is not None: + entry.updated_at = started_at + def mark_turn_active(self, session_key: str) -> Optional[str]: """Persist exact ownership of the agent turn running for *session_key*. The opaque token is returned to the caller and must be supplied to - :meth:`clear_turn_active`. Re-marking replaces the previous token so + :meth:`clear_turn_active`. Re-marking replaces the previous token so a stale asynchronous unwind cannot clear a newer turn. """ - from gateway.session import _now token = uuid.uuid4().hex with self._lock: entry = self._entry_locked(session_key) if entry is None: return None - now = _now() - candidate = entry.to_dict() - candidate["active_turn_token"] = token - candidate["active_turn_started_at"] = now.isoformat() - # Keeps the legacy 120s startup heuristic working for an older - # binary during a rolling downgrade/upgrade window. - candidate["updated_at"] = now.isoformat() - - # Persist before publishing in memory so a failed write cannot - # leak an unowned token through a later unrelated save. - self._save_entry(session_key, entry_data=candidate, lock_held=True) - entry.active_turn_token = token - entry.active_turn_started_at = now - entry.updated_at = now + self._set_turn_marker_locked(session_key, entry, token, _now()) return token def clear_turn_active(self, session_key: str, token: str) -> bool: - """Compare-and-swap clear an active-turn marker. - - Returns ``False`` when the entry disappeared or a newer turn owns it. - """ + """Compare-and-swap clear an active-turn marker; ``False`` when the + entry disappeared or a newer turn owns it.""" with self._lock: entry = self._entry_locked(session_key) if entry is None or entry.active_turn_token != token: return False - candidate = entry.to_dict() - candidate["active_turn_token"] = None - candidate["active_turn_started_at"] = None - - # Keep the live token until the clear is durable (retryable). - self._save_entry(session_key, entry_data=candidate, lock_held=True) - entry.active_turn_token = None - entry.active_turn_started_at = None + self._set_turn_marker_locked(session_key, entry, None, None) return True def recover_interrupted_turns( @@ -243,69 +272,58 @@ class SessionLifecycleMixin: Old/invalid markers are cleared without resuming; suspended sessions are never re-armed. Returns the number of newly promoted sessions. """ - from gateway.session import _now now = _now() max_age = timedelta(seconds=max(0, max_age_seconds)) promoted = 0 - changed = False - with self._lock: - self._ensure_loaded_locked() - for entry in self._entries.values(): - if not entry.active_turn_token: - continue + def _promote(entry: SessionEntry) -> bool: + nonlocal promoted + if not entry.active_turn_token: + return False + started_at = entry.active_turn_started_at + try: + marker_is_stale = ( + started_at is None + or (max_age_seconds > 0 and now - started_at > max_age) + ) + except TypeError: + # Mixed aware/naive timestamps: clear rather than risk an + # unsafe old resume. + marker_is_stale = True - started_at = entry.active_turn_started_at - try: - marker_is_stale = ( - started_at is None - or (max_age_seconds > 0 and now - started_at > max_age) - ) - except TypeError: - # Mixed aware/naive timestamps: clear rather than risk an - # unsafe old resume. - marker_is_stale = True - - if not marker_is_stale and not entry.suspended: - if entry.resume_pending: - # A drain-timeout marker is more specific; keep it. - if entry.last_resume_marked_at is None: - entry.last_resume_marked_at = now - else: - entry.resume_pending = True - entry.resume_reason = "restart_interrupted" - # Freshness starts at discovery, not turn start. + if not marker_is_stale and not entry.suspended: + if entry.resume_pending: + # A drain-timeout marker is more specific; keep it. + if entry.last_resume_marked_at is None: entry.last_resume_marked_at = now - promoted += 1 + else: + entry.resume_pending = True + entry.resume_reason = "restart_interrupted" + # Freshness starts at discovery, not turn start. + entry.last_resume_marked_at = now + promoted += 1 - entry.active_turn_token = None - entry.active_turn_started_at = None - changed = True - - if changed: - self._save() + entry.active_turn_token = None + entry.active_turn_started_at = None + return True + self._update_all_entries_locked(_promote) return promoted def discard_active_turn_markers(self) -> int: """Clear orphan turn markers after a verified clean shutdown.""" - cleared = 0 - with self._lock: - self._ensure_loaded_locked() - for entry in self._entries.values(): - if not entry.active_turn_token and entry.active_turn_started_at is None: - continue - entry.active_turn_token = None - entry.active_turn_started_at = None - cleared += 1 - if cleared: - self._save() - return cleared + def _discard(entry: SessionEntry) -> bool: + if not entry.active_turn_token and entry.active_turn_started_at is None: + return False + entry.active_turn_token = None + entry.active_turn_started_at = None + return True + + return self._update_all_entries_locked(_discard) def mark_resume_pending(self, session_key: str, reason: str = "restart_timeout") -> bool: """Mark a session resumable after a restart interruption (keeps the session_id/transcript, unlike ``suspend_session``). True if marked.""" - from gateway.session import _now def _apply(entry: SessionEntry): # Never override an explicit ``suspended`` (hard forced-wipe). if entry.suspended: @@ -335,22 +353,19 @@ class SessionLifecycleMixin: kept. The SQLite transcript stays; only the key -> session_id mapping is dropped. ``max_age_days <= 0`` disables. Returns the count removed. """ - from gateway.session import _now if max_age_days is None or max_age_days <= 0: return 0 cutoff = _now() - timedelta(days=max_age_days) - removed_keys: list[str] = [] with self._lock: self._ensure_loaded_locked() - for key, entry in list(self._entries.items()): - if entry.suspended: - continue + removed_keys = [ + key for key, entry in list(self._entries.items()) + if not entry.suspended # The callback is keyed by session_key, NOT session_id. - if self._has_active_processes_safe(entry.session_key, context="prune"): - continue - if entry.updated_at < cutoff: - removed_keys.append(key) + and not self._has_active_processes_safe(entry.session_key, context="prune") + and entry.updated_at < cutoff + ] for key in removed_keys: self._entries.pop(key, None) if removed_keys: @@ -367,19 +382,14 @@ class SessionLifecycleMixin: """Mark sessions active within *max_age_seconds* as ``resume_pending`` after a crash/fast restart (already-pending and suspended entries are skipped). Returns the number marked.""" - from gateway.session import _now cutoff = _now() - timedelta(seconds=max_age_seconds) - count = 0 - with self._lock: - self._ensure_loaded_locked() - for entry in self._entries.values(): - if entry.resume_pending: - continue - if not entry.suspended and entry.updated_at >= cutoff: - entry.resume_pending = True - entry.resume_reason = "restart_interrupted" - entry.last_resume_marked_at = _now() - count += 1 - if count: - self._save() - return count + + def _mark(entry: SessionEntry) -> bool: + if entry.resume_pending or entry.suspended or entry.updated_at < cutoff: + return False + entry.resume_pending = True + entry.resume_reason = "restart_interrupted" + entry.last_resume_marked_at = _now() + return True + + return self._update_all_entries_locked(_mark) diff --git a/gateway/session_persistence.py b/gateway/session_persistence.py index 9d1f044da1..96adecd8eb 100644 --- a/gateway/session_persistence.py +++ b/gateway/session_persistence.py @@ -10,6 +10,7 @@ from __future__ import annotations import logging import json import os +import tempfile import threading from pathlib import Path from typing import Any, Dict, Optional @@ -22,12 +23,32 @@ if TYPE_CHECKING: # Log-record parity with the origin module. logger = logging.getLogger("gateway.session") +# "No SessionDB pinned" sentinel: lets ``_db`` distinguish "resolve from the +# active scope" from a deliberate ``store._db = None`` (JSONL fallback). +_DB_UNPINNED = object() + +# Self-documenting sentinel written first into sessions.json; "_" keys are +# skipped on load. +_SESSIONS_JSON_README = ( + "LEGACY MIRROR of the gateway routing index (the primary copy " + "lives in the gateway_routing table in ~/.hermes/state.db). " + "Maps messaging session keys (agent:main::...) to " + "active session IDs. This is NOT the session list. ALL " + "sessions (CLI, TUI, and gateway) live in ~/.hermes/state.db " + "and are shown by `hermes sessions list` and `/sessions`. " + "Disable this file with `gateway.write_sessions_json: false` " + "in config.yaml." +) + + +def _is_live_system_guard(exc: BaseException) -> bool: + """Test-isolation guard: must stay a loud failure and is never cached.""" + return isinstance(exc, RuntimeError) and "live-system guard" in str(exc) + class SessionPersistenceMixin: - """SessionStore storage plumbing: per-profile SessionDB handle resolution and - the routing-index load/save paths (state.db gateway_routing primary, - sessions.json legacy mirror). - """ + """SessionStore storage plumbing: SessionDB handle resolution and the + routing-index load/save paths.""" def _open_session_db_for_active_scope(self, db_path: Optional[Path] = None): """SessionDB for the profile scope active on this task. @@ -41,29 +62,20 @@ class SessionPersistenceMixin: from hermes_state import _default_db_path, get_shared_session_db path = Path(db_path) if db_path is not None else Path(_default_db_path()) + def _open(): try: # Process-wide shared registry: one writer connection per path. return get_shared_session_db(path) except Exception as e: - if isinstance(e, RuntimeError) and "live-system guard" in str(e): - # Test-isolation guard: must stay a loud failure and is - # deliberately not cached so it fires again next attempt. - raise - print(f"[gateway] Warning: SQLite session store unavailable, falling back to JSONL: {e}") + if not _is_live_system_guard(e): + print(f"[gateway] Warning: SQLite session store unavailable, falling back to JSONL: {e}") raise - return self._db_handle_cache.get( - path, - _open, - non_cacheable=lambda exc: ( - isinstance(exc, RuntimeError) and "live-system guard" in str(exc) - ), - ) + return self._db_handle_cache.get(path, _open, non_cacheable=_is_live_system_guard) def _pinned_db(self): """Return the explicitly pinned DB (``store._db = x``), else ``_DB_UNPINNED``.""" - from gateway.session import _DB_UNPINNED return getattr(self, "_db_pinned", _DB_UNPINNED) @property @@ -75,7 +87,6 @@ class SessionPersistenceMixin: Unpinned, each read resolves the scope so a multiplexed profile's writes reach its own store. """ - from gateway.session import _DB_UNPINNED pinned = self._pinned_db() if pinned is not _DB_UNPINNED: return pinned @@ -96,7 +107,6 @@ class SessionPersistenceMixin: unrecovered. A pinned handle still wins; bare test instances lacking the handle cache report no DB. """ - from gateway.session import _DB_UNPINNED pinned = self._pinned_db() if pinned is not _DB_UNPINNED: return pinned @@ -124,11 +134,8 @@ class SessionPersistenceMixin: return profile def _profile_home_for_key(self, session_key: Optional[str]) -> Optional[Path]: - """HERMES_HOME of the profile that owns *session_key*, or None. - - None means only "no live home to point at" — no named owner, or the - owner's directory could not be resolved. - """ + """HERMES_HOME of the profile that owns *session_key*, or None (no named + owner, or the owner's directory could not be resolved).""" profile = self._named_profile_for_key(session_key) if profile is None: return None @@ -159,7 +166,6 @@ class SessionPersistenceMixin: write profile rows into the ROOT store until the stale-route self-heal drops a live conversation. The owning profile is encoded in the key. """ - from gateway.session import _DB_UNPINNED pinned = self._pinned_db() if pinned is not _DB_UNPINNED: return pinned @@ -249,6 +255,22 @@ class SessionPersistenceMixin: method = getattr(db, name, None) if db else None return method if callable(method) else None + def _load_routing_rows_locked(self) -> bool: + """Load state.db routing entries into ``_entries``; False when there is + no loader or the load failed (warned). Lock held.""" + loader = self._routing_db_method("load_gateway_routing_entries") + if loader is None: + return False + try: + for key, entry_json in loader(scope=self._routing_scope()).items(): + entry = self._routing_entry_from_json(key, entry_json) + if entry is not None: + self._entries[key] = entry + return True + except Exception as e: + logger.warning("gateway.session: state.db routing load failed: %s", e) + return False + @staticmethod def _routing_entry_from_json(key: str, entry_json: str) -> Optional[SessionEntry]: """Parse one gateway_routing row; None (with a warning) when invalid.""" @@ -273,28 +295,15 @@ class SessionPersistenceMixin: self.sessions_dir.mkdir(parents=True, exist_ok=True) - db_had_entries = False - db_load_succeeded = False - loader = self._routing_db_method("load_gateway_routing_entries") - if loader is not None: - try: - for key, entry_json in loader(scope=self._routing_scope()).items(): - entry = self._routing_entry_from_json(key, entry_json) - if entry is not None: - self._entries[key] = entry - db_had_entries = bool(self._entries) - db_load_succeeded = True - except Exception as e: - logger.warning("gateway.session: state.db routing load failed: %s", e) + db_load_succeeded = self._load_routing_rows_locked() + db_had_entries = db_load_succeeded and bool(self._entries) self._import_legacy_sessions_json(db_had_entries) self._loaded = True self._routing_db_loaded = db_load_succeeded self._routing_fallback_baseline = ( - None - if db_load_succeeded - else {key: entry.to_dict() for key, entry in self._entries.items()} + None if db_load_succeeded else self._entries_as_dicts() ) # A hard crash skips graceful shutdown and leaves sessions.json @@ -395,9 +404,7 @@ class SessionPersistenceMixin: logger.debug( "gateway.session: recovery lookup failed for stale " "sessions.json entry %r -> %s: %s", - key, - entry.session_id, - exc, + key, entry.session_id, exc, ) return None @@ -408,10 +415,7 @@ class SessionPersistenceMixin: logger.warning( "gateway.session: repointing stale sessions.json entry " "%r from ended %s (end_reason=%r) to recovered %s", - key, - entry.session_id, - row["end_reason"], - recovered_entry.session_id, + key, entry.session_id, row["end_reason"], recovered_entry.session_id, ) return recovered_entry @@ -433,6 +437,10 @@ class SessionPersistenceMixin: ) return "prune" + def _entries_as_dicts(self) -> Dict[str, Any]: + """Serializable snapshot of ``_entries``. Lock held.""" + return {key: entry.to_dict() for key, entry in self._entries.items()} + def _save(self) -> None: """Persist the routing index while the caller holds ``_lock``.""" data, generation = self._snapshot_routing_locked() @@ -463,7 +471,7 @@ class SessionPersistenceMixin: logger.warning("gateway.session: recovered state.db routing load failed: %s", exc) return - current = {key: entry.to_dict() for key, entry in self._entries.items()} + current = self._entries_as_dicts() for key, entry_json in durable.items(): durable_entry = self._routing_entry_from_json(key, entry_json) if durable_entry is None: @@ -486,10 +494,7 @@ class SessionPersistenceMixin: def _snapshot_routing_locked(self) -> tuple[Dict[str, Any], int]: """Capture immutable routing data and a monotonic generation.""" self._reconcile_recovered_routing_locked() - return ( - {key: entry.to_dict() for key, entry in self._entries.items()}, - self._next_routing_generation_locked(), - ) + return self._entries_as_dicts(), self._next_routing_generation_locked() def _persist_routing_data(self, data: Dict[str, Any], generation: int) -> None: """Serialize all whole-index writers through one durable write lock.""" @@ -507,10 +512,7 @@ class SessionPersistenceMixin: replacer = self._routing_db_method("replace_gateway_routing_entries") if replacer is not None: try: - replacer( - {k: json.dumps(v) for k, v in data.items()}, - scope=self._routing_scope(), - ) + replacer({k: json.dumps(v) for k, v in data.items()}, scope=self._routing_scope()) db_saved = True except Exception as exc: logger.warning("gateway.session: state.db routing save failed: %s", exc) @@ -531,36 +533,15 @@ class SessionPersistenceMixin: # This rewrite supersedes fast records at or below its # generation; newer ones stay for the next delayed full writer. if fast_persisted: - for key in [ - k for k, (rev, _) in fast_persisted.items() - if rev <= generation - ]: + for key in [k for k, (rev, _) in fast_persisted.items() if rev <= generation]: del fast_persisted[key] def _save_sessions_json(self, data: Dict[str, Any]) -> None: - """Write the legacy sessions.json mirror of the routing index.""" - import tempfile + """Write the legacy sessions.json mirror of the routing index (atomic + fsync).""" self.sessions_dir.mkdir(parents=True, exist_ok=True) sessions_file = self.sessions_dir / "sessions.json" - - # Self-documenting sentinel; "_" keys are skipped on load. Ordered - # first so it renders at the top of the file. - data = { - "_README": ( - "LEGACY MIRROR of the gateway routing index (the primary copy " - "lives in the gateway_routing table in ~/.hermes/state.db). " - "Maps messaging session keys (agent:main::...) to " - "active session IDs. This is NOT the session list. ALL " - "sessions (CLI, TUI, and gateway) live in ~/.hermes/state.db " - "and are shown by `hermes sessions list` and `/sessions`. " - "Disable this file with `gateway.write_sessions_json: false` " - "in config.yaml." - ), - **data, - } - fd, tmp_path = tempfile.mkstemp( - dir=str(self.sessions_dir), suffix=".tmp", prefix=".sessions_" - ) + data = {"_README": _SESSIONS_JSON_README, **data} + fd, tmp_path = tempfile.mkstemp(dir=str(self.sessions_dir), suffix=".tmp", prefix=".sessions_") try: with os.fdopen(fd, "w", encoding="utf-8") as f: json.dump(data, f, indent=2) @@ -606,19 +587,19 @@ class SessionPersistenceMixin: entry = self._entries.get(session_key) if entry is None: return None - serialized_entry = ( - dict(entry_data) if entry_data is not None else entry.to_dict() - ) + serialized_entry = dict(entry_data) if entry_data is not None else entry.to_dict() entry_json = json.dumps(serialized_entry) revision = self._next_routing_generation_locked() # The O(n) full snapshot is deferred to the fallback branch. return entry_json, revision, serialized_entry if entry_data is not None else None - if lock_held: - captured = _capture() - else: + def _locked(fn): + if lock_held: + return fn() with self._lock: - captured = _capture() + return fn() + + captured = _locked(_capture) if captured is None: return entry_json, revision, candidate_entry = captured @@ -643,13 +624,7 @@ class SessionPersistenceMixin: ) if candidate_entry is not None: # Full-snapshot fallback carrying the candidate transition. - def _snapshot() -> Dict[str, Any]: - return {key: current.to_dict() for key, current in self._entries.items()} - if lock_held: - fallback_data = _snapshot() - else: - with self._lock: - fallback_data = _snapshot() + fallback_data = _locked(self._entries_as_dicts) fallback_data[session_key] = candidate_entry self._persist_routing_data(fallback_data, revision) else: diff --git a/gateway/session_recovery.py b/gateway/session_recovery.py index 892c57b483..a3c01229bd 100644 --- a/gateway/session_recovery.py +++ b/gateway/session_recovery.py @@ -33,10 +33,7 @@ def _origin_json(source) -> Optional[str]: class SessionRecoveryMixin: - """SessionStore durable-row recovery: session-key generation, legacy Slack - key migration, rebuilding a routing entry from state.db, and the SQLite - side of routing transitions (promote/reopen/create/peer). - """ + """SessionStore durable-row recovery and the SQLite side of routing transitions.""" def _resolve_profile_for_key(self, source: Optional[SessionSource] = None) -> Optional[str]: """Profile namespace for session keys: None when multiplexing is off @@ -115,9 +112,7 @@ class SessionRecoveryMixin: """ if source.platform != Platform.SLACK or not source.scope_id: return None - return self._generate_session_key( - source, replace(source, scope_id=None, guild_id=None) - ) + return self._generate_session_key(source, replace(source, scope_id=None, guild_id=None)) def _claim_legacy_slack_key(self, legacy_key: Optional[str]) -> bool: """Atomically reserve one ambiguous legacy Slack key for migration.""" @@ -140,11 +135,7 @@ class SessionRecoveryMixin: names the same scope_id; rows without a parseable origin are rejected (an unattributable transcript is exactly the ambiguity to avoid). """ - if ( - source.platform != Platform.SLACK - or source.chat_type == "dm" - or not source.scope_id - ): + if source.platform != Platform.SLACK or source.chat_type == "dm" or not source.scope_id: return True try: origin = json.loads(recovered.get("origin_json") or "") @@ -163,6 +154,7 @@ class SessionRecoveryMixin: now: datetime, ) -> SessionEntry: from gateway.session import SessionEntry + def _ts(value, default: datetime) -> datetime: try: return datetime.fromtimestamp(float(value)) @@ -176,9 +168,7 @@ class SessionRecoveryMixin: updated_at = _ts(last_activity, created_at) if last_activity is not None else created_at had_activity = row.get("_has_messages") if had_activity is None: - had_activity = bool(row.get("message_count") or 0) or ( - last_activity is not None - ) + had_activity = bool(row.get("message_count") or 0) or last_activity is not None return SessionEntry( session_key=session_key, session_id=str(row["id"]), @@ -298,11 +288,7 @@ class SessionRecoveryMixin: raise_on_lookup_error=raise_on_lookup_error, ) migrated_legacy = False - if ( - not recovered - and legacy_key - and self._claim_legacy_slack_key(legacy_key) - ): + if not recovered and legacy_key and self._claim_legacy_slack_key(legacy_key): recovered = self._find_gateway_session_row( session_key=legacy_key, source=source, @@ -383,12 +369,11 @@ class SessionRecoveryMixin: thread_id=source.thread_id, ) try: - origin_json = _origin_json(source) recorder( session_id, **peer, display_name=display_name or source.chat_name, - origin_json=origin_json, + origin_json=_origin_json(source), include_compression_ancestors=include_compression_ancestors, ) except TypeError: @@ -483,7 +468,6 @@ class SessionRecoveryMixin: Identity (origin_json) and lineage (parent/_reset_from) land atomically in the INSERT so a crash right after cannot strand the row unroutable. """ - origin_json = _origin_json(origin) return { "session_id": session_id, "source": source_value, @@ -493,12 +477,10 @@ class SessionRecoveryMixin: "chat_type": origin.chat_type if origin else None, "thread_id": origin.thread_id if origin else None, "profile_name": origin.profile if origin else None, - "origin_json": origin_json, + "origin_json": _origin_json(origin), "display_name": display_name, "parent_session_id": parent_session_id, - "model_config": ( - {"_reset_from": parent_session_id} if parent_session_id else None - ), + "model_config": {"_reset_from": parent_session_id} if parent_session_id else None, } def _create_session_row(self, session_key, db_create_kwargs, origin, display_name, *, log) -> None: @@ -510,10 +492,7 @@ class SessionRecoveryMixin: try: self._db_for_key(session_key).create_session(**db_create_kwargs) self._record_gateway_session_peer( - db_create_kwargs["session_id"], - session_key, - origin, - display_name=display_name, + db_create_kwargs["session_id"], session_key, origin, display_name=display_name, ) except Exception as e: log(e) diff --git a/gateway/session_stall.py b/gateway/session_stall.py index 118981af6f..b2d06312b1 100644 --- a/gateway/session_stall.py +++ b/gateway/session_stall.py @@ -1,50 +1,36 @@ """Gateway session stall notification policy. -Consumes the shared activity observation contract from ``agent.session_activity`` -/ ``AIAgent.get_activity_summary()`` as the **single progress source**. This -module owns only the notify-once policy for "pending inbound + stale progress"; -it never derives a parallel progress clock from turn-start or inbound timestamps. - -Boundaries (keep separate): ``gateway/shutdown_watchdog.py`` is process / -event-loop liveness; ``gateway/delivery_ledger.py`` is outbound delivery -obligations. Pending inbound here is a stall *policy gate* (a queued follow-up -exists), not an obligation and not a progress timestamp. Timeout / kill / retry -policy stay in their own components. +Consumes ``AIAgent.get_activity_summary()`` as the **single progress source** and owns +only the notify-once policy for "pending inbound + stale progress"; it never derives a +parallel progress clock from turn-start or inbound timestamps. Pending inbound is a +stall *policy gate*, not a delivery obligation and not process liveness. """ from __future__ import annotations import math +import time from typing import Any, Mapping, Optional def should_emit_session_stall_notification( - *, - timeout_seconds: float, - idle_seconds: Optional[float], - has_pending_inbound: bool, + *, timeout_seconds: float, idle_seconds: Optional[float], has_pending_inbound: bool, already_notified: bool, ) -> bool: """Return True when a stall warning should be sent for this session.""" return ( - timeout_seconds > 0 - and has_pending_inbound - and not already_notified - and idle_seconds is not None - and idle_seconds >= timeout_seconds + timeout_seconds > 0 and has_pending_inbound and not already_notified + and idle_seconds is not None and idle_seconds >= timeout_seconds ) def should_clear_session_stall_notification( - *, - timeout_seconds: float, - idle_seconds: Optional[float], - has_pending_inbound: bool, + *, timeout_seconds: float, idle_seconds: Optional[float], has_pending_inbound: bool, ) -> bool: """Return True when a prior stall notice may be cleared (episode ended).""" if not has_pending_inbound or timeout_seconds <= 0: return True - # Unknown progress: hold the latch. Do not treat observation gaps as recovery. + # Unknown progress holds the latch: observation gaps are not recovery. return idle_seconds is not None and idle_seconds < timeout_seconds @@ -66,30 +52,18 @@ def _finite_float(value: Any) -> Optional[float]: def resolve_session_idle_seconds_from_activity( - activity: Optional[Mapping[str, Any]], - *, - now: Optional[float] = None, + activity: Optional[Mapping[str, Any]], *, now: Optional[float] = None, ) -> Optional[float]: - """Idle seconds from a shared activity snapshot only. - - Prefers ``seconds_since_activity`` when present and finite; otherwise derives - from ``last_activity_at`` / ``last_activity_ts``. Returns None when there is - no usable progress timestamp — callers must not fall back to turn-start or - pending-inbound clocks. - """ + """Idle seconds from a shared activity snapshot: a finite ``seconds_since_activity``, else + derived from ``last_activity_at`` / ``last_activity_ts``. None when there is no usable + progress timestamp — callers must not fall back to turn-start or inbound clocks.""" if not activity: return None - idle = _finite_float(activity.get("seconds_since_activity")) if idle is not None: return max(0.0, idle) - ts = activity.get("last_activity_at") when = _finite_float(activity.get("last_activity_ts") if ts is None else ts) if when is None: return None - if now is None: - import time as _time - - now = _time.time() - return max(0.0, float(now) - when) + return max(0.0, float(time.time() if now is None else now) - when) diff --git a/gateway/session_state.py b/gateway/session_state.py index 66ad914616..64b0605753 100644 --- a/gateway/session_state.py +++ b/gateway/session_state.py @@ -1,26 +1,9 @@ """Per-session gateway state consolidated into one container. -GatewayRunner historically carried ~19 separate ``Dict[str, ...]`` attributes -keyed by session_key, each with an ad-hoc lifecycle. That shape bred three -bug classes, all structurally closed here: - -1. Boundary drift — hand-copied pop-lists at conversation boundaries went - stale when a dict was added. Now one ``ConversationState.clear()``. -2. Turn-release drift — ad-hoc ``del self._running_agents[key]`` sites popped - different subsets of the turn dicts. Now ``TurnState.clear()``. -3. Wholesale-reset races — lazy-init ``self._x = {}`` replaced the ENTIRE dict, - discarding concurrent sessions' entries. Resets now touch one field of - one ``SessionState``. - -Scopes follow where each dict was CLEARED: ``turn`` resets at the end of every -running turn; ``conversation`` at conversation boundaries (/new, /resume, -auto-reset, expiry, compression-exhausted reset); ``persistent`` fields have -their own lifecycles and ``run_generation`` is monotonic and NEVER reset. - -Entries in ``GatewayRunner._sessions`` are never evicted (matching the old -dicts, which also leaked empty/stale entries for dead sessions). Eviction of -fully-default SessionStates is possible follow-up work. -""" +Replaces ~19 separate session_key-keyed dicts on GatewayRunner that bred boundary drift, +turn-release drift and wholesale-reset races. Scopes follow where each dict was CLEARED: +``turn`` at the end of every turn; ``conversation`` at conversation boundaries (/new, +/resume, auto-reset, expiry); ``persistent`` fields have their own lifecycles.""" from __future__ import annotations @@ -28,91 +11,67 @@ from collections.abc import MutableMapping from dataclasses import dataclass, field from typing import Any, Callable, Dict, Iterator, List, NamedTuple, Optional, Tuple -# Presence-sensitive sentinel: /fast stores "priority" or None (explicit -# normal), so key PRESENCE — not value truthiness — decides whether the -# override applies. ``_UNSET_TIER`` means "no override recorded". +# Presence-sensitive sentinel: /fast stores "priority" or None (explicit normal), so key +# PRESENCE — not value truthiness — decides whether the override applies. _UNSET_TIER = object() SERVICE_TIER_UNSET = _UNSET_TIER # public alias @dataclass class TurnState: - """State scoped to one running gateway turn. + """State scoped to one running gateway turn. ``lease_token`` / ``lease_generation`` + are NOT touched by ``clear()``: ``_release_turn_lease`` owns them (release exactly once).""" - ``clear()`` runs at every site that ends a running turn. ``lease_token`` / - ``lease_generation`` are deliberately NOT cleared by it — they are owned by - ``_release_turn_lease``, which must release the registry lease exactly once - per acquiring turn. - """ - - # Running AIAgent instance (or _AGENT_PENDING_SENTINEL); None = idle. - agent: Any = None + agent: Any = None # running AIAgent (or _AGENT_PENDING_SENTINEL); None = idle started_ts: float = 0.0 # 0.0 = not running lease: Any = None # cross-process active-session slot lease busy_ack_ts: float = 0.0 # debounce; 0.0 = never acked - # Held turn-lease token + the run generation that acquired it. The pair - # replaces the old (session_key, generation)-keyed dict so a stale unwind - # can never free a newer turn's lease: release/rebind only match when the - # generation is current. + # Held turn-lease token + the generation that acquired it: release/rebind only match + # when the generation is current, so a stale unwind can never free a newer turn's lease. lease_token: Any = None lease_generation: Optional[int] = None def clear(self) -> None: - """Reset the per-turn slot (agent / start ts / lease / busy-ack). - - The caller pops ``lease`` first so it can call ``lease.release()``. - """ - self.agent = None - self.started_ts = 0.0 - self.lease = None - self.busy_ack_ts = 0.0 + """Reset the per-turn slot. The caller pops ``lease`` first to release it.""" + self.agent = self.lease = None + self.started_ts = self.busy_ack_ts = 0.0 @dataclass class ConversationState: """State scoped to one conversation (survives turns, not boundaries).""" - # /model per-session override (model/provider/api_key/base_url/api_mode). - model_override: Optional[Dict[str, Any]] = None + model_override: Optional[Dict[str, Any]] = None # /model per-session override one_turn_restore: Optional[Dict[str, Any]] = None # /model --once snapshot reasoning_override: Optional[Dict[str, Any]] = None # /reasoning override - # /fast per-session override: "priority" or None; _UNSET_TIER = absent. - service_tier_override: Any = _UNSET_TIER + service_tier_override: Any = _UNSET_TIER # /fast: "priority" or None; _UNSET_TIER = absent last_resolved_model: str = "" # last successfully-resolved non-empty model - queued_events: List[Any] = field(default_factory=list) # /queue overflow FIFO (adapter slot holds the head) + queued_events: List[Any] = field(default_factory=list) # /queue overflow FIFO (head in adapter) sidecar_notes: List[str] = field(default_factory=list) # one-shot must-deliver notes - ephemeral_pin: Optional[Tuple[Any, ...]] = None # pinned session-context bytes: (change_key, text) + ephemeral_pin: Optional[Tuple[Any, ...]] = None # pinned session-context (change_key, text) vc_last: Optional[str] = None # last voice-channel context delivered def clear(self) -> None: - """Reset every field to its default — adding a field here means every - conversation boundary clears it automatically.""" + """Reset every field to its default, so new fields are cleared automatically.""" self.__dict__.update(ConversationState().__dict__) @dataclass class PersistentState: - """State with its own lifecycle — NOT cleared wholesale by turn or boundary - resets (approvals/update prompts ARE cleared by the boundary *security* - funnel, but individually).""" + """State with its own lifecycle — NOT cleared wholesale by turn or boundary resets + (approvals/update prompts ARE cleared, individually, by the boundary security funnel).""" approvals: Optional[Dict[str, Any]] = None # {"command": ..., "pattern_key": ...} update_prompt_pending: bool = False # /update prompt awaiting a reply native_image_paths: List[str] = field(default_factory=list) # consumed one-shot - # Legacy runner-level pending message text (write-mostly; flushed to disk on - # shutdown). Distinct from the adapter-level ``_pending_messages`` - # (Dict[str, MessageEvent]) in gateway/base.py, which shares the old name. + # Legacy runner-level pending text (flushed on shutdown); distinct from gateway/base.py's + # adapter-level ``_pending_messages``. pending_command_text: Optional[str] = None # Monotonic run-generation counter. NEVER reset: stale-run detection depends on it. run_generation: int = 0 - # Consecutive session-hygiene compression failures. The in-agent compressor's - # own timeout ladder is unreachable from the gateway (hygiene builds a FRESH - # AIAgent per run and bind_session_state() zeroes that counter), so the streak - # lives here and lets hygiene escalate its cooldown. Reset on a successful - # compression only. PROCESS-LOCAL by design: no disk flush, so a restart drops - # escalation to rung 1 while the DB-backed deadline survives; gateway.run - # mirrors it to the DB keyed by session_key (not session_id) so it also holds - # across compaction ROTATION, where the sid changes but the chat does not. + # Consecutive hygiene compression failures, so hygiene can escalate its cooldown (the + # in-agent compressor's ladder is unreachable: hygiene builds a FRESH AIAgent per run). + # Reset only on success. PROCESS-LOCAL; gateway.run mirrors it to the DB by session_key. hygiene_failure_streak: int = 0 @@ -125,20 +84,13 @@ class SessionState: persistent: PersistentState = field(default_factory=PersistentState) -# --------------------------------------------------------------------------- -# Legacy dict-view adapters. -# -# Dozens of tests construct bare runners (object.__new__) and read/write the -# old dict attributes directly (``runner._running_agents = {}``, ``assert key -# in runner._pending_approvals``). Each view is a LIVE MutableMapping over one -# SessionState field across all sessions. Production code in gateway/run.py -# uses ``self._session_state(key)..``. -# --------------------------------------------------------------------------- +# --- Legacy dict-view adapters --------------------------------------------- +# Dozens of tests read/write the old dict attributes directly (``runner._running_agents = +# {}``). Each view is a LIVE MutableMapping over one SessionState field across sessions. class _FieldSpec(NamedTuple): """One legacy dict: scope attr, field name, default factory, presence test.""" - scope: str name: str default: Callable[[], Any] @@ -146,8 +98,7 @@ class _FieldSpec(NamedTuple): def _spec(scope: str, name: str, default: Any) -> _FieldSpec: - """``default`` is either a type (presence = truthiness) or a sentinel value - such as ``None`` / ``_UNSET_TIER`` (presence = ``is not`` sentinel).""" + """``default`` is a type (presence = truthiness) or a sentinel (presence = ``is not``).""" if isinstance(default, type): return _FieldSpec(scope, name, default, bool) return _FieldSpec(scope, name, lambda: default, lambda v: v is not default) @@ -167,16 +118,11 @@ class _RunnerView(MutableMapping): def __len__(self) -> int: return sum(1 for _ in self) - # Mapping doesn't provide __eq__; tests compare against plain dicts. - def __eq__(self, other: object) -> bool: + def __eq__(self, other: object) -> bool: # Mapping has no __eq__; tests compare to dicts if isinstance(other, (dict, MutableMapping)): return dict(self.items()) == dict(other) return NotImplemented - def __ne__(self, other: object) -> bool: - result = self.__eq__(other) - return NotImplemented if result is NotImplemented else not result - class SessionFieldView(_RunnerView): """Live dict-like view of one SessionState field across sessions.""" @@ -193,29 +139,30 @@ class SessionFieldView(_RunnerView): def _set(self, state: SessionState, value: Any) -> None: setattr(getattr(state, self._spec.scope), self._spec.name, value) - def __getitem__(self, key: str) -> Any: + def _present(self, key: Any) -> Optional[SessionState]: + """The session state for ``key`` if its field is present, else None.""" state = self._sessions().get(key) - if state is None or not self._spec.is_present(value := self._value(state)): + return state if state is not None and self._spec.is_present(self._value(state)) else None + + def _held(self, key: str) -> SessionState: + if (state := self._present(key)) is None: raise KeyError(key) - return value + return state + + def __getitem__(self, key: str) -> Any: + return self._value(self._held(key)) def __setitem__(self, key: str, value: Any) -> None: self._set(self._runner._session_state(key), value) def __delitem__(self, key: str) -> None: - state = self._sessions().get(key) - if state is None or not self._spec.is_present(self._value(state)): - raise KeyError(key) - self._set(state, self._spec.default()) + self._set(self._held(key), self._spec.default()) def __iter__(self) -> Iterator[str]: - for key, state in list(self._sessions().items()): - if self._spec.is_present(self._value(state)): - yield key + return (k for k in list(self._sessions()) if self._present(k) is not None) def __contains__(self, key: object) -> bool: - state = self._sessions().get(key) # type: ignore[arg-type] - return state is not None and self._spec.is_present(self._value(state)) + return self._present(key) is not None def clear(self) -> None: # avoid MutableMapping's popitem loop for state in list(self._sessions().values()): @@ -226,26 +173,23 @@ class SessionFieldView(_RunnerView): class TurnLeaseTokenView(_RunnerView): - """Legacy view of ``_turn_lease_tokens``: keyed by (session_key, generation). - - The pair lives on ``TurnState.lease_token`` / ``lease_generation``; the lease - registry serializes acquisition per session, so at most one held token - exists per session key and the single slot equals the old tuple-keyed dict. - """ + """Legacy view of ``_turn_lease_tokens``, keyed by (session_key, generation). The lease + registry serializes acquisition per session, so the single ``TurnState`` slot per key + equals the old tuple-keyed dict.""" __slots__ = () - def _held(self, key: Any) -> Tuple[Any, SessionState]: - """Return (session_key, state) for a currently-held (key, gen) or raise KeyError.""" + def _held(self, key: Any) -> TurnState: + """TurnState for a currently-held (session_key, generation) or raise KeyError.""" if not isinstance(key, tuple) or len(key) != 2: raise KeyError(key) state = self._sessions().get(key[0]) if state is None or state.turn.lease_token is None or state.turn.lease_generation != key[1]: raise KeyError(key) - return key[0], state + return state.turn def __getitem__(self, key: Any) -> Any: - return self._held(key)[1].turn.lease_token + return self._held(key).lease_token def __setitem__(self, key: Any, value: Any) -> None: if not isinstance(key, tuple) or len(key) != 2: @@ -254,13 +198,12 @@ class TurnLeaseTokenView(_RunnerView): turn.lease_token, turn.lease_generation = value, key[1] def __delitem__(self, key: Any) -> None: - turn = self._held(key)[1].turn + turn = self._held(key) turn.lease_token = turn.lease_generation = None def __iter__(self) -> Iterator[Tuple[str, Any]]: - for key, state in list(self._sessions().items()): - if state.turn.lease_token is not None: - yield (key, state.turn.lease_generation) + return ((k, s.turn.lease_generation) for k, s in list(self._sessions().items()) + if s.turn.lease_token is not None) def clear(self) -> None: # avoid MutableMapping's popitem loop for key in list(self): @@ -291,19 +234,12 @@ LEGACY_FIELD_SPECS: Dict[str, _FieldSpec] = { def _legacy_property(make_view: Callable[[Any], MutableMapping], doc: str) -> property: - """Dict-shaped @property over a live view. - - Getter returns the view; setter accepts a plain dict (the ubiquitous test - pattern ``runner._X = {...}``), resetting the field on every known session - and then applying the given entries; ``del runner._X`` (older tests - simulating a runner without the attribute) means "no entries". - """ - + """Dict-shaped @property over a live view. The setter takes a plain dict (test pattern + ``runner._X = {...}``): reset the field on every session, then apply the entries.""" def fset(self: Any, mapping: Optional[Dict[Any, Any]]) -> None: view = make_view(self) view.clear() - for key, value in (mapping or {}).items(): - view[key] = value + view.update(mapping or {}) return property(make_view, fset, lambda self: make_view(self).clear(), doc=doc) diff --git a/gateway/session_transcript.py b/gateway/session_transcript.py index e2f4430451..14af7b9ac7 100644 --- a/gateway/session_transcript.py +++ b/gateway/session_transcript.py @@ -18,6 +18,14 @@ if TYPE_CHECKING: logger = logging.getLogger("gateway.session") +class TranscriptReadError(RuntimeError): + """Raised when persisted history cannot be read safely.""" + + def __init__(self, session_id: str) -> None: + self.session_id = session_id + super().__init__(f"transcript read failed for session {session_id}") + + def _plain_text(content) -> str: """Text of a message content (str or text-part list); "" for anything else.""" if isinstance(content, list): @@ -26,11 +34,22 @@ def _plain_text(content) -> str: return content if isinstance(content, str) else "" +def _spool_dropped(session_id: str, message: Dict[str, Any]): + """Spool one evicted/undeliverable message to disk (same machinery as the + shutdown flush, so it is replayed after DB recovery); path or None.""" + try: + from gateway.shutdown_flush import spool_dropped_transcript_message + + return spool_dropped_transcript_message(session_id, message) + except Exception: + return None + + class SessionTranscriptMixin: """SessionStore transcript I/O: SQLite append with a per-session retry queue, - compression-reroute following, FTS corruption recovery, - rewrite/rewind/load. - """ + compression-reroute following, FTS corruption recovery, rewrite/rewind/load.""" + + _MAX_PENDING_PER_SESSION = 200 # in-memory pending messages per session (DB broken) def _compression_tip_for_session_id(self, session_id: Optional[str]) -> Optional[str]: """Latest compression continuation for *session_id* (heals a mapping @@ -62,8 +81,7 @@ class SessionTranscriptMixin: return False logger.info( "SessionStore healed compressed session mapping: %s -> %s", - entry.session_id, - canonical_session_id, + entry.session_id, canonical_session_id, ) entry.session_id = canonical_session_id return True @@ -91,11 +109,7 @@ class SessionTranscriptMixin: return entry if entry.session_id != expected_session_id: return None - if not self._heal_compression_tip_locked( - entry, - expected_session_id, - target_session_id, - ): + if not self._heal_compression_tip_locked(entry, expected_session_id, target_session_id): return None # Bookkeeping, not user activity: leave ``updated_at`` alone. self._save() @@ -110,9 +124,7 @@ class SessionTranscriptMixin: if not self._db_for_session_id(session_id) or skip_db: return with self._get_transcript_drain_lock(): - self._append_to_transcript_serialized( - self._follow_reroutes(session_id), message - ) + self._append_to_transcript_serialized(self._follow_reroutes(session_id), message) def _follow_reroutes(self, session_id: str) -> str: """Follow the compression reroute chain (cycle-guarded).""" @@ -123,25 +135,12 @@ class SessionTranscriptMixin: session_id = reroutes[session_id] return session_id - def _spool_dropped(self, session_id: str, message: Dict[str, Any]): - """Spool one evicted/undeliverable message to disk; path or None.""" - try: - from gateway.shutdown_flush import spool_dropped_transcript_message - - return spool_dropped_transcript_message(session_id, message) - except Exception: - return None - def _enqueue_transcript_message(self, session_id: str, message: Dict[str, Any]) -> list: - """Queue *message* (retry lock held); evicts + spools the oldest past the cap. - - Spooling uses the same machinery as shutdown flush so the message is - replayed after DB recovery instead of being lost. - """ + """Queue *message* (retry lock held); evicts + spools the oldest past the cap.""" pending = self._dirty_transcripts.setdefault(session_id, []) pending.append(dict(message)) if len(pending) > self._MAX_PENDING_PER_SESSION: - spool_path = self._spool_dropped(session_id, pending.pop(0)) + spool_path = _spool_dropped(session_id, pending.pop(0)) if spool_path is not None: self._lazy("_spooled_drop_sessions", set).add(session_id) logger.warning( @@ -240,8 +239,7 @@ class SessionTranscriptMixin: previous_failures = self._transcript_append_failures.pop(queue_session_id, 0) if previous_failures: self._transcript_append_failures[child_id] = max( - previous_failures, - self._transcript_append_failures.get(child_id, 0), + previous_failures, self._transcript_append_failures.get(child_id, 0), ) self._transcript_reroutes[session_id] = child_id return pending @@ -372,10 +370,7 @@ class SessionTranscriptMixin: from gateway.shutdown_flush import drain_transcript_spool _replayed, remaining = drain_transcript_spool( - session_id, - lambda message: self._append_transcript_message( - session_id, message - ), + session_id, lambda message: self._append_transcript_message(session_id, message), ) if not remaining: spooled_sessions.discard(session_id) @@ -415,8 +410,6 @@ class SessionTranscriptMixin: display_metadata=message.get("display_metadata"), ) - _MAX_PENDING_PER_SESSION = 200 - @staticmethod def _is_fts_corruption_error(exc: Exception) -> bool: """True only when the failure is provably scoped to the FTS index. @@ -426,16 +419,13 @@ class SessionTranscriptMixin: ``SessionDB._is_fts_write_corruption_error``) may authorize the one-shot rebuild-and-retry. Everything else takes the retry path. """ - text = str(exc).lower() - if "messages_fts" in text: + if "messages_fts" in str(exc).lower(): return True import sqlite3 from hermes_state import SessionDB - if isinstance(exc, sqlite3.DatabaseError): - return SessionDB._is_fts_write_corruption_error(exc) - return False + return isinstance(exc, sqlite3.DatabaseError) and SessionDB._is_fts_write_corruption_error(exc) def _rebuild_fts_once(self) -> bool: """Attempt FTS5 ``rebuild`` once per store lifetime; True if any index was rebuilt.""" @@ -462,10 +452,7 @@ class SessionTranscriptMixin: logger.warning("Session DB FTS rebuild failed: %s", exc) return False if rebuilt: - logger.warning( - "Rebuilt %d Session DB FTS index(es) after append corruption", - rebuilt, - ) + logger.warning("Rebuilt %d Session DB FTS index(es) after append corruption", rebuilt) return rebuilt > 0 def _clear_dirty_transcript(self, session_id: str) -> None: @@ -482,9 +469,7 @@ class SessionTranscriptMixin: if not db: return False try: - return db.has_platform_message_id( - session_id, platform_message_id - ) + return db.has_platform_message_id(session_id, platform_message_id) except Exception: logger.debug("has_platform_message_id lookup failed", exc_info=True) return False @@ -529,7 +514,6 @@ class SessionTranscriptMixin: (compression rotation), then the durable compression tip — otherwise the transcript "vanishes" while every message sits under the child. """ - from gateway.session import TranscriptReadError if not self._db_for_session_id(session_id): return [] session_id = self._follow_reroutes(session_id) @@ -552,9 +536,7 @@ class SessionTranscriptMixin: logger.error( "Transcript read failed for session %s; refusing to treat the " "conversation as empty: %s", - session_id, - e, - exc_info=True, + session_id, e, exc_info=True, ) raise TranscriptReadError(session_id) from e @@ -577,8 +559,7 @@ class SessionTranscriptMixin: if not db: return None with self._get_transcript_drain_lock(): - if n < 1: - n = 1 + n = max(n, 1) from agent.context_compressor import ( retryable_user_text, split_user_originated_turn, @@ -587,10 +568,7 @@ class SessionTranscriptMixin: try: expected_active_ids = db.get_active_message_ids(session_id) - durable = db.get_messages_as_conversation( - session_id, - include_row_ids=True, - ) + durable = db.get_messages_as_conversation(session_id, include_row_ids=True) user_indices = [ index for index, message in enumerate(durable)