From d7bdf2788d3a8e3272b419e61e5f5fbbf77133e9 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:15:42 -0700 Subject: [PATCH] refactor(gateway/session): split SessionStore into persistence/recovery/lifecycle/transcript mixins by call-graph cohesion; compact wire helpers --- gateway/session.py | 2207 +------------------------------- gateway/session_lifecycle.py | 385 ++++++ gateway/session_persistence.py | 656 ++++++++++ gateway/session_recovery.py | 519 ++++++++ gateway/session_transcript.py | 641 ++++++++++ 5 files changed, 2264 insertions(+), 2144 deletions(-) create mode 100644 gateway/session_lifecycle.py create mode 100644 gateway/session_persistence.py create mode 100644 gateway/session_recovery.py create mode 100644 gateway/session_transcript.py diff --git a/gateway/session.py b/gateway/session.py index 744da51322..1c536e6226 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -17,7 +17,7 @@ import threading import uuid from pathlib import Path from datetime import datetime, timedelta -from dataclasses import dataclass, field, fields, replace +from dataclasses import dataclass, field, fields from typing import Dict, List, Optional, Any logger = logging.getLogger(__name__) @@ -112,9 +112,10 @@ from .whatsapp_identity import ( canonical_whatsapp_identifier, normalize_whatsapp_identifier, # noqa: F401 - re-exported for gateway.session callers ) -from utils import atomic_replace -from agent.turn_context import extract_api_content_sidecar -import contextlib +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. @@ -218,57 +219,51 @@ class SessionSource: parts.append(f"thread: {self.thread_id}") return ", ".join(parts) + # Wire layout (order matters for byte-stable JSON): always-present fields, + # then truthy-only optionals around the dual-written scope pair. + _ALWAYS_FIELDS = ("chat_id", "chat_name", "chat_type", "user_id", "user_name", "thread_id", "chat_topic") + _OPTIONAL_PRE_SCOPE = ("user_id_alt", "chat_id_alt") + _OPTIONAL_POST_SCOPE = ("parent_chat_id", "message_id", "profile") + _OPTIONAL_TAIL = ("auto_thread_initial_name", "prospective_thread_id") + def to_dict(self) -> Dict[str, Any]: - d = { - "platform": self.platform.value, - "chat_id": self.chat_id, - "chat_name": self.chat_name, - "chat_type": self.chat_type, - "user_id": self.user_id, - "user_name": self.user_name, - "thread_id": self.thread_id, - "chat_topic": self.chat_topic, - } - def _optional(*names: str) -> None: + d = {"platform": self.platform.value} + 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 - _optional("user_id_alt", "chat_id_alt") + _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 - _optional("parent_chat_id", "message_id", "profile") + _optional(self._OPTIONAL_POST_SCOPE) if self.auto_thread_created: d["auto_thread_created"] = True - _optional("auto_thread_initial_name", "prospective_thread_id") + _optional(self._OPTIONAL_TAIL) return d @classmethod def from_dict(cls, data: Dict[str, Any]) -> "SessionSource": + plain = { + name: data.get(name) + for name in cls._ALWAYS_FIELDS[1:] + cls._OPTIONAL_PRE_SCOPE + cls._OPTIONAL_POST_SCOPE + cls._OPTIONAL_TAIL + if name != "chat_type" + } return cls( platform=Platform(data["platform"]), chat_id=str(data["chat_id"]), - chat_name=data.get("chat_name"), chat_type=data.get("chat_type", "dm"), - user_id=data.get("user_id"), - user_name=data.get("user_name"), - thread_id=data.get("thread_id"), - chat_topic=data.get("chat_topic"), - user_id_alt=data.get("user_id_alt"), - chat_id_alt=data.get("chat_id_alt"), scope_id=data.get("scope_id", data.get("guild_id")), - parent_chat_id=data.get("parent_chat_id"), - message_id=data.get("message_id"), - profile=data.get("profile"), auto_thread_created=bool(data.get("auto_thread_created", False)), - auto_thread_initial_name=data.get("auto_thread_initial_name"), - prospective_thread_id=data.get("prospective_thread_id"), + **plain, ) - + @dataclass @@ -278,13 +273,13 @@ 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 updated_at: Optional[datetime] = None - + def to_dict(self) -> Dict[str, Any]: return { "source": self.source.to_dict(), @@ -340,11 +335,9 @@ def _slack_tools_loaded() -> bool: try: from hermes_cli.config import load_config from hermes_cli.tools_config import _get_platform_tools - cfg = load_config() # include_default_mcp_servers defaults True so a default-enabled Slack # MCP server counts too. - enabled = _get_platform_tools(cfg, "slack") - return "slack" in enabled + return "slack" in _get_platform_tools(load_config(), "slack") except Exception: return False @@ -360,8 +353,7 @@ def _discord_tools_loaded() -> bool: if not (get_secret("DISCORD_BOT_TOKEN", "") or "").strip(): return False - cfg = load_config() - enabled = _get_platform_tools(cfg, "discord", include_default_mcp_servers=False) + enabled = _get_platform_tools(load_config(), "discord", include_default_mcp_servers=False) return "discord" in enabled or "discord_admin" in enabled except Exception: return False @@ -649,10 +641,10 @@ class SessionEntry: session_id: str created_at: datetime updated_at: datetime - + # Origin metadata for delivery routing origin: Optional[SessionSource] = None - + # Display metadata display_name: Optional[str] = None platform: Optional[Platform] = None @@ -661,7 +653,7 @@ class SessionEntry: # Small, JSON-serializable per-entry state (e.g. Slack thread watermarks); # persisted in the routing index. metadata: Dict[str, Any] = field(default_factory=dict) - + # Token tracking input_tokens: int = 0 output_tokens: int = 0 @@ -670,10 +662,10 @@ 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 - + # Set when a session was created because the previous one expired; # consumed once by the message handler to inject a notice into context was_auto_reset: bool = False @@ -748,13 +740,10 @@ class SessionEntry: if self.origin: result["origin"] = self.origin.to_dict() return result - + @classmethod def from_dict(cls, data: Dict[str, Any]) -> "SessionEntry": - origin = None - if "origin" in data and isinstance(data["origin"], dict): - origin = SessionSource.from_dict(data["origin"]) - + origin = SessionSource.from_dict(data["origin"]) if isinstance(data.get("origin"), dict) else None platform = None if data.get("platform"): try: @@ -777,13 +766,9 @@ class SessionEntry: # 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): - raise ValueError( - "Invalid session_id: potential directory traversal detected" - ) + raise ValueError("Invalid session_id: potential directory traversal detected") if _is_path_unsafe(session_key, strict=False): - raise ValueError( - "Invalid session_key: potential directory traversal detected" - ) + 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} @@ -825,10 +810,8 @@ def build_channel_continuity_note( """ if source.platform not in (Platform.SLACK, Platform.DISCORD): return None - if not getattr(entry, "reset_had_activity", False): - return None - prev = getattr(entry, "prev_session_id", None) - if not prev: + prev = entry.prev_session_id + if not entry.reset_had_activity or not prev: return None where = "thread" if source.thread_id else "channel" @@ -978,7 +961,12 @@ class AsyncSessionStore: _DB_UNPINNED = object() -class SessionStore: +class SessionStore( + SessionPersistenceMixin, + SessionRecoveryMixin, + SessionLifecycleMixin, + SessionTranscriptMixin, +): """Session storage/retrieval: SQLite (SessionDB) for metadata and transcripts, legacy JSONL fallback when SQLite is unavailable.""" @@ -1018,37 +1006,32 @@ class SessionStore: self._fts_rebuild_attempted = False self._has_active_processes_fn = has_active_processes_fn # Keep the legacy sessions.json mirror (disable via gateway.write_sessions_json). - self._write_sessions_json = bool( - getattr(config, "write_sessions_json", True) - ) + self._write_sessions_json = bool(getattr(config, "write_sessions_json", True)) - # SQLite handles are cached per resolved path and looked up through - # the ``_db`` property rather than bound once here: a multiplexed - # gateway serves every profile from ONE process, and a handle bound in - # __init__ would be frozen to the root home so every profile's rows - # land in the root state.db. Priming the current scope's handle below - # keeps startup diagnostics (live-DB guard, JSONL warning) at - # construction time. + # SQLite handles are cached per resolved path and resolved through the + # ``_db`` property, never bound once here: a multiplexed gateway serves + # every profile from ONE process and a handle frozen to the root home + # would land every profile's rows in the root state.db. Priming the + # current scope below keeps startup diagnostics at construction time. self._db_pinned = _DB_UNPINNED self._db_handles: Dict[Path, Any] = {} self._db_handles_lock = threading.Lock() - # profile name -> its HERMES_HOME; memoized so the per-key store - # lookup is a dict hit, not a profile-directory stat per append. + # profile name -> HERMES_HOME; memoized so per-key store lookup is a + # dict hit, not a profile-directory stat per append. self._profile_home_cache: Dict[str, Optional[Path]] = {} - # session_id -> owning routing key, for ids whose ownership is proven - # but not yet published in ``_entries`` (compression continuation: - # the child row is written before its reroute is published). + # session_id -> owning routing key for ids whose ownership is proven + # but not yet published in ``_entries`` (compression child row is + # written before its reroute is published). self._session_owner_hints: Dict[str, str] = {} from gateway.session_db_recovery import RecoverableHandleCache self._db_handle_cache = RecoverableHandleCache( - handles=self._db_handles, - lock=self._db_handles_lock, + handles=self._db_handles, lock=self._db_handles_lock, ) # The routing index is one process-wide structure keyed by - # ``agent::…``, so it needs exactly one home for its lifetime; - # capture the gateway's own home at startup (before any profile scope - # exists) — see ``_routing_db``. + # ``agent::…`` and needs exactly one home for its lifetime: + # the gateway's own, captured before any profile scope exists + # (see ``_routing_db``). try: from hermes_constants import get_hermes_home @@ -1070,203 +1053,6 @@ class SessionStore: setattr(self, name, value) return value - def _open_session_db_for_active_scope(self, db_path: Optional[Path] = None): - """Return the SessionDB for the profile scope active on this task. - - ``db_path`` pins the store explicitly; otherwise ``_default_db_path()`` - follows the context-local HERMES_HOME installed by - ``_profile_runtime_scope`` (resolving per call, not once in - ``__init__``, is what lets multiplexed profiles reach their own store). - Handles are cached per resolved path; failed opens enter a bounded - backoff during which callers keep using the JSONL fallback. - """ - 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}") - raise - - return self._db_handle_cache.get( - path, - _open, - non_cacheable=lambda exc: ( - isinstance(exc, RuntimeError) and "live-system guard" in str(exc) - ), - ) - - def _pinned_db(self): - """Return the explicitly pinned DB (``store._db = x``), else ``_DB_UNPINNED``.""" - return getattr(self, "_db_pinned", _DB_UNPINNED) - - @property - def _db(self): - """The SessionDB for the active profile scope, or a pinned override. - - Assigning ``store._db`` pins that value for every subsequent read - (tests install a fake or disable the DB with ``store._db = None``). - Unpinned, each read resolves the scope so a multiplexed profile's - writes reach its own store. - """ - pinned = self._pinned_db() - if pinned is not _DB_UNPINNED: - return pinned - return self._open_session_db_for_active_scope() - - @_db.setter - def _db(self, value) -> None: - self._db_pinned = value - - @property - def _routing_db(self): - """The one store that owns the routing index, whatever scope is active. - - ``_entries`` is a single flat dict holding every profile's keys, so it - must persist to a single file (``_routing_home``), not whichever - profile happens to be scoped — otherwise a rewrite during one - profile's turn and the unscoped startup load see different copies, - and crash markers written under a secondary profile go unrecovered. - A pinned handle still wins. Bare test instances lacking the handle - cache report no DB. - """ - pinned = self._pinned_db() - if pinned is not _DB_UNPINNED: - return pinned - home = getattr(self, "_routing_home", None) - try: - if home is None: - return self._db - return self._open_session_db_for_active_scope(db_path=home / "state.db") - except Exception: - return None - - def _named_profile_for_key(self, session_key: Optional[str]) -> Optional[str]: - """The non-default profile that owns *session_key*, or None. - - None means the ambient store is authoritative (multiplexing off, or - legacy ``agent:main`` namespace). It deliberately does NOT cover "that - profile has no directory" — ownership and resolvability are separate - questions that ``_db_for_key`` answers separately. - """ - if not getattr(self.config, "multiplex_profiles", False): - return None - profile = self._profile_from_session_key(session_key) - if not profile or profile == "default": - return None - 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. - """ - profile = self._named_profile_for_key(session_key) - if profile is None: - return None - cache = self._profile_home_cache - if profile in cache: - return cache[profile] - home: Optional[Path] = None - try: - from hermes_cli.profiles import get_profile_dir, profile_exists - - if profile_exists(profile): - home = Path(get_profile_dir(profile)) - except Exception as exc: - logger.debug("Could not resolve profile home for %r: %s", session_key, exc) - home = None - # Only hits are memoized: a profile directory can be provisioned - # *after* startup (enrollment bridge), and a cached miss would pin - # that profile's rows to the ambient store for the process lifetime. - if home is not None: - cache[profile] = home - return home - - def _db_for_key(self, session_key: Optional[str]): - """The SessionDB holding *session_key*'s rows, whatever scope is active. - - ``_db`` follows the ambient HERMES_HOME, which only the inbound - message path installs; background work (e.g. the expiry watcher) - runs unscoped over every profile's keys and would otherwise write - profile rows into the ROOT store, drifting from the real row until - the stale-route self-heal drops a live conversation. The owning - profile is encoded in the key, so derive the store from it. - """ - pinned = self._pinned_db() - if pinned is not _DB_UNPINNED: - return pinned - profile = self._named_profile_for_key(session_key) - if profile is None: - return self._db - home = self._profile_home_for_key(session_key) - if home is None: - # Named owner we cannot resolve (not provisioned yet, or lookup - # failed). Falling back to the ambient store would split ONE - # session identity across two physical stores — fail closed; - # callers already handle a missing DB. - logger.warning( - "gateway.session: profile %r has no resolvable home (key %r); " - "refusing to fall back to the ambient store", - profile, session_key, - ) - return None - try: - return self._open_session_db_for_active_scope(db_path=home / "state.db") - except Exception: - # Same contract as ``_db``: a failed open degrades to JSONL fallback. - return None - - def _owner_key_for_session_id(self, session_id: Optional[str]) -> Optional[str]: - """The routing key that owns *session_id*, or None. - - The published index is authoritative; ``_session_owner_hints`` covers - the window where ownership is proven but routing not yet published. - Deliberately lock-free: several callers already hold ``_lock``. - """ - if not session_id: - return None - try: - for entry in list(self._entries.values()): - if entry.session_id == session_id: - return entry.session_key - except Exception: - pass # bare stores / foreign entry objects in suites - return (getattr(self, "_session_owner_hints", None) or {}).get(session_id) - - def _db_for_session_id(self, session_id: Optional[str]): - """The SessionDB holding *session_id*'s row (owner recovered from the - index or a pre-published hint; unknown ids fall back to the ambient store).""" - if not session_id: - return self._db - return self._db_for_key(self._owner_key_for_session_id(session_id)) - - def close_all_db_handles(self) -> None: - """Close every SessionDB handle this store opened, one per resolved path. - - Closing just ``store._db`` at shutdown would strand every secondary - profile's handle with its WAL lock held (restart flows then hit - 'database is locked'). Handles are drained under the lock but closed - outside it so concurrent resolvers never wait on N ``close()`` calls. - A pinned handle is deliberately not closed — the pinner owns it. - """ - def _close(db) -> None: - # Shared instances no-op on close(); release the refcount instead. - from hermes_state import release_or_close - try: - release_or_close(db) - except Exception as exc: - logger.debug("SessionDB close error during handle sweep: %s", exc) - - self._db_handle_cache.close_all(_close) def _has_active_processes_safe(self, session_key: str, *, context: str) -> bool: """Return whether a session has active work, failing closed on registry errors.""" @@ -1282,947 +1068,17 @@ class SessionStore: exc, ) return True - - def _ensure_loaded(self) -> None: - """Load sessions index from disk if not already loaded.""" - with self._lock: - self._ensure_loaded_locked() - def _entry_locked(self, session_key: str) -> Optional[SessionEntry]: - """Load the index and return the entry for *session_key*. Lock held.""" - self._ensure_loaded_locked() - return self._entries.get(session_key) - def _routing_scope(self) -> str: - """Namespace for this store's gateway_routing rows: the resolved - sessions_dir, so stores with different dirs never share entries.""" - try: - return str(Path(self.sessions_dir).resolve()) - except Exception: - return str(self.sessions_dir) - def _routing_db_method(self, name: str): - """Bound ``_routing_db.`` if the handle exists and has it, else None.""" - db = self._routing_db - method = getattr(db, name, None) if db else None - return method if callable(method) else None - @staticmethod - def _routing_entry_from_json(key: str, entry_json: str) -> Optional[SessionEntry]: - """Parse one gateway_routing row; None (with a warning) when invalid.""" - try: - entry_data = json.loads(entry_json) - if isinstance(entry_data, dict): - return SessionEntry.from_dict(entry_data) - except (ValueError, KeyError, TypeError) as e: - logger.warning("Skipping invalid routing entry %r: %s", key, e) - return None - def _ensure_loaded_locked(self) -> None: - """Load the routing index. Must be called with self._lock held. - state.db ``gateway_routing`` is primary; sessions.json is the legacy - import for keys the DB lacks (persisted to the DB on the next _save). - """ - if self._loaded: - self._reconcile_recovered_routing_locked() - return - 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) - # Legacy import: sessions.json only fills keys the DB lacks. - sessions_file = self.sessions_dir / "sessions.json" - if sessions_file.exists(): - try: - with open(sessions_file, "r", encoding="utf-8") as f: - data = json.load(f) - imported = 0 - for key, entry_data in data.items(): - # "_"-prefixed keys are sentinels (e.g. "_README"), not entries. - if key.startswith("_"): - continue - if key in self._entries: - continue - # A non-dict entry (corrupt file) must not abort the - # whole load. - if not isinstance(entry_data, dict): - logger.warning( - "Skipping invalid session entry %r: " - "expected dict, got %s", - key, type(entry_data).__name__, - ) - continue - try: - self._entries[key] = SessionEntry.from_dict(entry_data) - imported += 1 - except (ValueError, KeyError, TypeError) as e: - logger.warning("Skipping invalid session entry %r: %s", key, e) - if imported and db_had_entries: - logger.info( - "gateway.session: imported %d legacy sessions.json " - "entr%s missing from state.db routing table", - imported, "y" if imported == 1 else "ies", - ) - except Exception as e: - print(f"[gateway] Warning: Failed to load sessions: {e}") - 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()} - ) - # A hard crash skips graceful shutdown and leaves sessions.json - # pointing at ended sessions; self-heal before the first message. - self._prune_stale_sessions_locked() - - def _prune_stale_sessions_locked(self) -> None: - """Remove routing entries whose session has ended in state.db (startup, lock held). - - Stale == ``end_reason IS NOT NULL``. Rows absent from the DB are kept; - a ``None`` DB handle is a no-op; DB errors are non-fatal. - """ - if not self._entries: - return - - stale_keys: list = [] - recovered_keys = 0 - try: - for key, entry in self._entries.items(): - # Ask the store that owns the key, not the ambient handle, or a - # live secondary-profile session gets pruned on the root copy. - db = self._db_for_key(key) - if db is None: - continue - row = db.get_session(entry.session_id) - if row is not None and row.get("end_reason") is not None: - recovered_entry = None - if entry.origin is not None: - try: - recovered_entry = self._recover_session_from_db( - session_key=key, - source=entry.origin, - now=_now(), - raise_on_lookup_error=True, - ) - except Exception as exc: - # Indeterminate: keep the only routing handle. - logger.debug( - "gateway.session: recovery lookup failed for stale " - "sessions.json entry %r -> %s: %s", - key, - entry.session_id, - exc, - ) - continue - - # Compression-ended parent with a newer live child for the - # same peer: repoint instead of dropping, or queued/ - # resume-pending work vanishes until the next message. - if recovered_entry is not None and recovered_entry.session_id != entry.session_id: - 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, - ) - self._entries[key] = recovered_entry - recovered_keys += 1 - continue - - # Same-id recovery == successful resume: keep the ORIGINAL - # entry object (the recovered one is rebuilt minimal and - # would drop counters, model_override, resume markers, - # metadata). Nothing changes, so no save. - if recovered_entry is not None: - logger.info( - "gateway.session: reopened ended session %s for " - "sessions.json entry %r (end_reason=%r); keeping route", - entry.session_id, key, row["end_reason"], - ) - continue - - logger.warning( - "gateway.session: pruning stale sessions.json entry " - "%r -> %s (end_reason=%r); left by a crashed gateway", - key, entry.session_id, row["end_reason"], - ) - stale_keys.append(key) - except Exception as exc: - logger.warning( - "gateway.session: stale-entry pruning skipped due to DB error: %s", - exc, - ) - return - - for key in stale_keys: - del self._entries[key] - - if stale_keys or recovered_keys: - self._save() - - def _save(self) -> None: - """Persist the routing index while the caller holds ``_lock``.""" - data, generation = self._snapshot_routing_locked() - self._persist_routing_data(data, generation) - - def _next_routing_generation_locked(self) -> int: - """Bump and return the shared routing counter. Caller holds ``_lock``. - - Full snapshots AND single-entry fast saves MUST allocate from this one - counter: the stale-write protection is a total order over - serialization times and silently breaks otherwise. - """ - self._routing_generation = getattr(self, "_routing_generation", 0) + 1 - return self._routing_generation - - def _reconcile_recovered_routing_locked(self) -> None: - """Merge authoritative rows after a fallback-only startup load.""" - baseline = getattr(self, "_routing_fallback_baseline", None) - if getattr(self, "_routing_db_loaded", False) or baseline is None: - return - - loader = self._routing_db_method("load_gateway_routing_entries") - if loader is None: - return - try: - durable = loader(scope=self._routing_scope()) - except Exception as exc: - 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()} - for key, entry_json in durable.items(): - durable_entry = self._routing_entry_from_json(key, entry_json) - if durable_entry is None: - continue - - if key not in baseline: - # A key created while on fallback wins over a DB-only key; - # otherwise restore the authoritative row that fallback never saw. - self._entries.setdefault(key, durable_entry) - elif key not in current: - # The key was loaded from fallback and deliberately removed. - continue - elif current[key] == baseline[key]: - # Unchanged fallback data yields to the authoritative DB copy. - self._entries[key] = durable_entry - - self._routing_db_loaded = True - self._routing_fallback_baseline = None - - 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(), - ) - - def _persist_routing_data(self, data: Dict[str, Any], generation: int) -> None: - """Serialize all whole-index writers through one durable write lock.""" - with self._lazy("_save_lock", threading.Lock): - if generation <= getattr(self, "_persisted_routing_generation", 0): - return - # Fold in fast upserts numbered above this snapshot: they were - # serialized after us and a delayed full rewrite must not regress them. - fast_persisted = getattr(self, "_fast_persisted_entries", None) - if fast_persisted: - for key, (revision, entry_json) in fast_persisted.items(): - if revision > generation: - data[key] = json.loads(entry_json) - db_saved = False - 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(), - ) - db_saved = True - except Exception as exc: - logger.warning("gateway.session: state.db routing save failed: %s", exc) - if getattr(self, "_write_sessions_json", True) or not db_saved: - try: - self._save_sessions_json(data) - except Exception as exc: - if not db_saved: - raise - # state.db is authoritative. A failed legacy mirror must not - # report the already-committed primary write as failed. - logger.warning( - "gateway.session: sessions.json mirror save failed " - "after state.db commit: %s", - exc, - ) - self._persisted_routing_generation = generation - # 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 - ]: - 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 - 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_" - ) - try: - with os.fdopen(fd, "w", encoding="utf-8") as f: - json.dump(data, f, indent=2) - f.flush() - os.fsync(f.fileno()) - atomic_replace(tmp_path, sessions_file) - except BaseException: - try: - os.unlink(tmp_path) - except OSError as e: - logger.debug("Could not remove temp file %s: %s", tmp_path, e) - raise - - def _save_entries(self) -> None: - """Snapshot latest state under ``_lock`` and persist after releasing it.""" - with self._lock: - data, generation = self._snapshot_routing_locked() - self._persist_routing_data(data, generation) - - def _save_entry( - self, - session_key: str, - *, - entry_data: Optional[Dict[str, Any]] = None, - lock_held: bool = False, - ) -> None: - """Persist ONE routing entry via UPSERT — the per-turn fast path. - - A full index rewrite re-serializes every entry and fsyncs a multi-MB - sessions.json (~50ms at ~1100 keys, twice per turn); a single-row - UPSERT takes well under a millisecond. Invariants: - - - The key -> session_id mapping never changes here; structural - transitions (create/recover/reset/switch/prune/heal) use the full - rewrite, which also refreshes the legacy sessions.json mirror. The - mirror may lag in metadata only; state.db stays primary. - - Ordering: the entry is serialized under ``_lock`` with a revision - from the shared routing generation counter, so a higher number - always means same-or-newer data. Under ``_save_lock`` the upsert is - skipped if a full snapshot or a fast save of this key with a higher - number already persisted. The reverse (delayed full rewrite after a - later fast save) is handled in ``_persist_routing_data``. - - No DB, or a failed upsert, falls back to the full rewrite so - DB-less installs keep sessions.json durable every turn. - - ``entry_data`` persists a candidate before it is published to the live - entry (failure-atomic metadata transitions); the full-save fallback - carries the same candidate. - """ - def _capture() -> Optional[tuple[str, int, Optional[Dict[str, Any]]]]: - 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() - ) - 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: - with self._lock: - captured = _capture() - if captured is None: - return - entry_json, revision, candidate_entry = captured - saver = self._routing_db_method("save_gateway_routing_entry") - if saver is not None: - try: - with self._lazy("_save_lock", threading.Lock): - if getattr(self, "_persisted_routing_generation", 0) >= revision: - return - fast_persisted = self._lazy("_fast_persisted_entries", dict) - persisted = fast_persisted.get(session_key) - if persisted is not None and persisted[0] >= revision: - return - saver(session_key, entry_json, scope=self._routing_scope()) - fast_persisted[session_key] = (revision, entry_json) - return - except Exception as exc: - logger.warning( - "gateway.session: single-entry routing save failed for %r " - "(%s); falling back to full index rewrite", - session_key, exc, - ) - 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[session_key] = candidate_entry - self._persist_routing_data(fallback_data, revision) - else: - self._save_entries() - - def _resolve_profile_for_key(self, source: Optional[SessionSource] = None) -> Optional[str]: - """Profile namespace for session keys: None when multiplexing is off - (legacy ``agent:main``), else ``source.profile`` or the active profile.""" - if not getattr(self.config, "multiplex_profiles", False): - return None - if source is not None and source.profile: - return source.profile - try: - from hermes_cli.profiles import get_active_profile_name - return get_active_profile_name() or "default" - except Exception: - return None - - @staticmethod - def _profile_from_session_key(session_key: Optional[str]) -> Optional[str]: - """Extract the profile namespace encoded in a gateway session key.""" - if not session_key: - return None - parts = str(session_key).split(":") - if len(parts) < 2 or parts[0] != "agent": - return None - namespace = parts[1] or "main" - return "default" if namespace == "main" else namespace - - @staticmethod - def _active_profile_name() -> str: - try: - from hermes_cli.profiles import get_active_profile_name - return get_active_profile_name() or "default" - except Exception: - return "default" - - def _recovered_row_allowed_for_active_profile( - self, - *, - requested_session_key: str, - recovered: Dict[str, Any], - ) -> bool: - """Prevent a gateway from reviving another profile's row. - - Single-profile: the row's namespace must match the ACTIVE profile. - Multiplexed: it must match the namespace of the requested key (the - active profile is meaningless there). Keyless rows stay adoptable. - """ - recovered_key = str(recovered.get("session_key") or "") - if not recovered_key or recovered_key == requested_session_key: - return True - - recovered_profile = self._profile_from_session_key(recovered_key) - if recovered_profile is None: - return True - - if getattr(self.config, "multiplex_profiles", False): - requested_profile = self._profile_from_session_key(requested_session_key) - return requested_profile is None or recovered_profile == requested_profile - - return recovered_profile == self._active_profile_name() - - def _generate_session_key(self, source: SessionSource, key_source: Optional[SessionSource] = None) -> str: - """Session key for *source* (profile resolved from *source*, key built - from *key_source* when given).""" - return build_session_key( - key_source if key_source is not None else source, - group_sessions_per_user=getattr(self.config, "group_sessions_per_user", True), - thread_sessions_per_user=getattr(self.config, "thread_sessions_per_user", False), - profile=self._resolve_profile_for_key(source), - ) - - def _legacy_slack_session_key(self, source: SessionSource) -> Optional[str]: - """Pre-workspace Slack key for an explicitly scoped source. - - Deliberately Slack-only; an unscoped Slack session may be claimed by - only one workspace because its old key cannot distinguish teams. - """ - 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) - ) - - def _claim_legacy_slack_key(self, legacy_key: Optional[str]) -> bool: - """Atomically reserve one ambiguous legacy Slack key for migration.""" - if not legacy_key: - return False - with self._lazy("_legacy_slack_claim_lock", threading.Lock): - claimed = self._lazy("_claimed_legacy_slack_keys", set) - if legacy_key in claimed: - return False - claimed.add(legacy_key) - return True - - @staticmethod - def _recovered_row_matches_source_scope( - recovered: Dict[str, Any], source: SessionSource - ) -> bool: - """Reject recovered rows whose recorded origin belongs to another workspace. - - A workspace-scoped Slack lookup adopts a row only if its origin_json - 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 - ): - return True - try: - origin = json.loads(recovered.get("origin_json") or "") - except (TypeError, ValueError): - return False - if not isinstance(origin, dict): - return False - return origin.get("scope_id", origin.get("guild_id")) == source.scope_id - - def _create_entry_from_recovered_row( - self, - *, - row: Dict[str, Any], - session_key: str, - source: SessionSource, - now: datetime, - ) -> SessionEntry: - def _ts(value, default: datetime) -> datetime: - try: - return datetime.fromtimestamp(float(value)) - except (TypeError, ValueError, OSError): - return default - - # An invalid durable timestamp must look old, never freshly active. - created_at = _ts(row.get("started_at"), datetime.fromtimestamp(0)) - # The finder already returns durable recency; no extra round-trip. - last_activity = row.get("last_activity_at") - 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 - ) - return SessionEntry( - session_key=session_key, - session_id=str(row["id"]), - created_at=created_at, - updated_at=updated_at, - origin=source, - display_name=source.chat_name, - platform=source.platform, - chat_type=source.chat_type, - reset_had_activity=bool(had_activity), - ) - - def _find_gateway_session_row( - self, - *, - session_key: str, - source: SessionSource, - allow_peer_fallback: bool, - raise_on_lookup_error: bool = False, - ) -> Optional[Dict[str, Any]]: - """Query one durable gateway session row. - - Scoped Slack lookups disable SessionDB's platform/chat/user fallback: - that tuple does not contain a workspace id and could therefore revive - another team's session. The caller performs one explicit exact lookup - of the old unscoped key instead. - """ - db = self._db_for_key(session_key) - finder = getattr(db, "find_latest_gateway_session_for_peer", None) if db else None - if not callable(finder): - return None - try: - return finder( - source=source.platform.value, - user_id=source.user_id, - session_key=session_key, - chat_id=source.chat_id if allow_peer_fallback else None, - chat_type=source.chat_type if allow_peer_fallback else None, - thread_id=source.thread_id, - ) - except Exception as exc: - logger.debug("Gateway session DB recovery failed for %s: %s", session_key, exc) - if raise_on_lookup_error: - raise - return None - - def _recover_session_from_db( - self, - *, - session_key: str, - source: SessionSource, - now: datetime, - raise_on_lookup_error: bool = False, - ) -> Optional[SessionEntry]: - """Rebuild a missing session-key mapping from durable state.db data. - - Returns ``None`` when no row is recoverable, or when the recovered - session is already overdue under the reset policy — the row is then - durably promoted to a reset boundary instead of resurrected. - """ - entry, migrated_legacy = self._query_recoverable_row( - session_key=session_key, - source=source, - now=now, - raise_on_lookup_error=raise_on_lookup_error, - ) - if entry is None: - return None - reset_reason = self._should_reset(entry, source) - if reset_reason: - self._promote_session_reset( - session_key, entry.session_id, reset_reason, - log=lambda exc: logger.debug( - "Gateway recovered-session reset promotion failed for %s: %s", - session_key, exc, - ), - ) - return None - self._reopen_session_row(session_key, entry.session_id) - if migrated_legacy: - self._record_gateway_session_peer( - entry.session_id, session_key, source, display_name=entry.display_name, - ) - return entry - - def _query_recoverable_session(self, *, session_key, source, now): - """DB-only half of _recover_session_from_db (no lock needed). - - Returns a SessionEntry or None. Caller assigns _entries[key] under - lock. The row is NOT reopened here: the caller evaluates reset policy - first (an agent_close/ws_orphan row may need promotion to a real reset - boundary instead). - """ - entry, migrated_legacy = self._query_recoverable_row( - session_key=session_key, source=source, now=now, - ) - if entry is not None and migrated_legacy: - self._record_gateway_session_peer( - entry.session_id, session_key, source, display_name=entry.display_name, - ) - return entry - - def _query_recoverable_row( - self, *, session_key, source, now, raise_on_lookup_error=False, - ) -> tuple[Optional[SessionEntry], bool]: - """Find and gate a recoverable row -> (entry or None, migrated_legacy). - - The legacy (pre-workspace) Slack key fallback lives here: exact-key - lookup, claimed once per process; ``migrated_legacy`` tells the caller - to rewrite the peer row to the scoped key. - """ - legacy_key = self._legacy_slack_session_key(source) - recovered = self._find_gateway_session_row( - session_key=session_key, - source=source, - allow_peer_fallback=legacy_key is None, - 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) - ): - recovered = self._find_gateway_session_row( - session_key=legacy_key, - source=source, - allow_peer_fallback=False, - raise_on_lookup_error=raise_on_lookup_error, - ) - migrated_legacy = bool(recovered) - if not isinstance(recovered, dict): - return None, False - if not self._recovered_row_matches_source_scope(recovered, source): - return None, False - if not self._recovered_row_allowed_for_active_profile( - requested_session_key=session_key, - recovered=recovered, - ): - logger.warning( - "Gateway session DB recovery ignored %s for %s because " - "the row belongs to a different profile", - recovered.get("session_key"), - session_key, - ) - return None, False - entry = self._create_entry_from_recovered_row( - row=recovered, session_key=session_key, source=source, now=now, - ) - return entry, migrated_legacy - - def _promote_session_reset(self, session_key: str, session_id: str, reason: str, *, log) -> None: - """End *session_id* with *reason* via ``promote_to_session_reset``. - - Promote (not plain ``end_session``): a row already ended with a - recoverable accidental reason (agent_close / ws_orphan_reap) must be - upgraded to the explicit boundary, or stale-route recovery resurrects - it over the reset. Falls back to ``end_session`` on old SessionDBs. - ``log(exc)`` reports failures (each caller has its own message). - """ - try: - db = self._db_for_key(session_key) - promote = getattr(db, "promote_to_session_reset", None) - if callable(promote): - promote(session_id, reason) - else: - db.end_session(session_id, reason) - except Exception as exc: - log(exc) - - def _reopen_session_row(self, session_key: str, session_id: str, *, log_prefix: str = "") -> None: - """Best-effort ``reopen_session``; failures are debug-logged only.""" - try: - self._db_for_key(session_key).reopen_session(session_id) - except Exception as exc: - if log_prefix: - logger.debug("%s: %s", log_prefix, exc) - else: - logger.debug("Gateway session DB reopen failed for %s: %s", session_key, exc) - - def _record_gateway_session_peer( - self, - session_id: str, - session_key: str, - source: Optional[SessionSource], - display_name: Optional[str] = None, - include_compression_ancestors: bool = False, - ) -> None: - """Persist the routing peer for an existing gateway session row.""" - db = self._db_for_key(session_key) - if not db or not source: - return - recorder = getattr(db, "record_gateway_session_peer", None) - if not callable(recorder): - return - peer = dict( - source=source.platform.value, - user_id=source.user_id, - session_key=session_key, - chat_id=source.chat_id, - chat_type=source.chat_type, - thread_id=source.thread_id, - ) - try: - origin_json = None - with contextlib.suppress(Exception): - origin_json = json.dumps(source.to_dict()) - recorder( - session_id, - **peer, - display_name=display_name or source.chat_name, - origin_json=origin_json, - include_compression_ancestors=include_compression_ancestors, - ) - except TypeError: - # Older SessionDB without display_name/origin_json kwargs. - try: - recorder(session_id, **peer) - except Exception as exc: - logger.debug("Gateway session peer record failed for %s: %s", session_key, exc) - except Exception as exc: - logger.debug("Gateway session peer record failed for %s: %s", session_key, exc) - - def set_expiry_finalized( - self, entry: SessionEntry, *, clear_model_override: bool = True - ) -> None: - """Mark a session entry expiry-finalized in memory, sessions.json, AND state.db. - - Single write-path for the expiry watcher so the durable flag survives - sessions.json loss. ``clear_model_override=False`` = flag only. - """ - with self._lock: - entry.expiry_finalized = True - if clear_model_override: - # Finalization is a conversation boundary: drop the persisted - # /model override so a later message cannot rehydrate it. - entry.model_override = None - self._save() - # 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) - 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.""" - 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, - ) - if now.hour < policy.at_hour: - today_reset -= timedelta(days=1) - if updated_at < today_reset: - return "daily" - return None - - def _is_session_expired(self, entry: SessionEntry) -> bool: - """Whether the entry's reset policy has expired it (entry alone, no source). - - Used by the background expiry watcher. Sessions with active - background processes are never considered expired. - """ - 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, - ) - return self._policy_reset_reason(policy, entry.updated_at) is not None - - def is_session_finalizable(self, entry: SessionEntry) -> bool: - """True if the expiry watcher will *ever* finalize this session. - - A ``mode == "none"`` session never expires, so the agent-cache idle - sweep must reap its agent itself instead of deferring to the watcher - (deferring would pin the agent for the gateway's lifetime). Policy - resolution errors count as "not finalizable" (sweep reaps — safe). - """ - try: - policy = self.config.get_reset_policy( - platform=entry.platform, - session_type=entry.chat_type, - ) - return policy.mode != "none" - except Exception: - return False - - def _is_session_ended_in_db(self, session_id: str) -> bool: - """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 - session ended while the gateway stays alive. Store resolved from the - row's owning profile, not the ambient scope. - """ - db = self._db_for_session_id(session_id) - if not db or not session_id: - return False - try: - row = db.get_session(session_id) - except Exception: - return False - 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. - """ - 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 - ) - return self._policy_reset_reason(policy, entry.updated_at) - - def _compression_tip_for_session_id(self, session_id: Optional[str]) -> Optional[str]: - """Latest compression continuation for *session_id* (heals a mapping - left pointing at a compressed parent by a restart or failed send).""" - if not session_id: - return session_id - db = self._db_for_session_id(session_id) - if db is None: - return session_id - try: - return db.get_compression_tip(session_id) or session_id - except Exception: - logger.debug("Compression-tip lookup failed for session %s", session_id, exc_info=True) - return session_id - - def _heal_compression_tip_locked( - self, - entry: "SessionEntry", - original_session_id: Optional[str], - canonical_session_id: Optional[str], - ) -> bool: - """Rewrite *entry* to the compression continuation if stale. Lock held.""" - if ( - not original_session_id - or not canonical_session_id - or entry.session_id != original_session_id - or canonical_session_id == original_session_id - ): - return False - logger.info( - "SessionStore healed compressed session mapping: %s -> %s", - entry.session_id, - canonical_session_id, - ) - entry.session_id = canonical_session_id - return True def has_any_sessions(self) -> bool: """Whether any session has ever been created (across all platforms). @@ -2292,106 +1148,6 @@ class SessionStore: with inflight_lock: self._inflight_sessions.pop(session_key, None) - def _adopt_legacy_slack_entry(self, source: SessionSource, session_key: str) -> None: - """One-time migration of pre-workspace-scope Slack keys. - - MOVE (not copy) the legacy entry so a second workspace with identical - Slack ids cannot attach to the same transcript. Adopt when the legacy - origin names the same workspace; a scope-less DM is claimed once by - the first workspace; a scope-less channel/group is refused (channel - ids collide across workspaces). - """ - legacy_key = self._legacy_slack_session_key(source) - if not legacy_key: - return - migrated: Optional[SessionEntry] = None - with self._lock: - self._ensure_loaded_locked() - legacy_entry = self._entries.get(legacy_key) - if session_key not in self._entries and legacy_entry is not None: - origin_scope = getattr(legacy_entry.origin, "scope_id", None) - if origin_scope is not None: - adopt = origin_scope == source.scope_id - else: - adopt = source.chat_type == "dm" - if adopt and self._claim_legacy_slack_key(legacy_key): - migrated = self._entries.pop(legacy_key) - migrated.session_key = session_key - migrated.origin = source - migrated.platform = source.platform - migrated.chat_type = source.chat_type - self._entries[session_key] = migrated - if migrated is not None: - self._save_entries() - self._record_gateway_session_peer( - migrated.session_id, session_key, source, display_name=migrated.display_name, - ) - - def _route_reset_reason( - self, entry: SessionEntry, source: SessionSource, now: datetime - ) -> Optional[str]: - """Reset decision for an existing route (no lock; DB/config I/O). - - ``suspended`` always resets. Otherwise the reset policy decides; a - still-pending resume marker is additionally freshness-gated — but - ``session_reset.mode: none`` (user opted out of ALL automatic resets) - makes an expired marker fall through to a normal resume, never a - silent fresh session. - """ - 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, - ) - if policy.mode == "none": - return None - window = auto_continue_freshness_window() - ref_time = entry.last_resume_marked_at or entry.updated_at - if window > 0 and (now - ref_time).total_seconds() > window: - return "resume_pending_expired" - return None - - def _finish_route_transition( - self, - session_key: str, - *, - end_session_id: Optional[str], - end_reason: str, - create_kwargs: Optional[Dict[str, Any]], - origin: Optional[SessionSource], - display_name: Optional[str], - during: str = "", - ) -> None: - """SQLite side of a routing transition, outside ``_lock``. - - Promotes the predecessor row to an explicit reset boundary (with the - specific reason so state.db is auditable, e.g. ``resume_pending_expired`` - vs a plain ``session_reset``), then INSERTs the new row + routing peer. - Both are best-effort: failures are warned and self-healed by the next - per-turn peer refresh. - """ - if self._db_for_key(session_key) and end_session_id: - self._promote_session_reset( - session_key, end_session_id, end_reason, - log=lambda e: logger.warning( - "Failed to end predecessor session row %s for %s%s: %s — " - "the old row remains open and may win restart recovery " - "until the next successful peer refresh", - end_session_id, session_key, during, e, - ), - ) - if self._db_for_key(session_key) and create_kwargs: - self._create_session_row( - session_key, create_kwargs, origin, display_name, - log=lambda e: logger.warning( - "Failed to create session row %s for %s%s: %s — deferring " - "to the self-healing peer refresh on the next turn", - create_kwargs.get("session_id"), session_key, during, e, - ), - ) def _get_or_create_session_impl( self, @@ -2563,54 +1319,6 @@ class SessionStore: ) return entry - @staticmethod - def _session_create_kwargs( - *, session_id, session_key, origin, source_value, display_name, parent_session_id, - ) -> Dict[str, Any]: - """kwargs for ``SessionDB.create_session``. - - 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 = None - if origin is not None: - try: - origin_json = json.dumps(origin.to_dict()) - except Exception: - origin_json = None - return { - "session_id": session_id, - "source": source_value, - "user_id": origin.user_id if origin else None, - "session_key": session_key, - "chat_id": origin.chat_id if origin else None, - "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, - "display_name": display_name, - "parent_session_id": parent_session_id, - "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: - """INSERT a session row and record its routing peer; ``log(exc)`` on failure. - - A failed create is a routing hazard (visible warning), but the row is - self-healed with full identity by the next per-turn peer refresh. - """ - 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, - ) - except Exception as e: - log(e) def update_session( self, @@ -2645,15 +1353,6 @@ class SessionStore: display_name=peer_display_name, ) - def _update_entry(self, session_key: str, mutate) -> bool: - """Apply ``mutate(entry)`` under ``_lock`` and full-save; False when the - entry is missing or *mutate* returned False (nothing to persist).""" - with self._lock: - entry = self._entry_locked(session_key) - if entry is None or mutate(entry) is False: - return False - self._save() - return True def get_session_metadata(self, session_key: str, key: str, default: Any = None) -> Any: """Return a metadata value stored on a live session entry.""" @@ -2688,203 +1387,6 @@ class SessionStore: entry = self._entry_locked(session_key) return dict(entry.model_override) if entry and entry.model_override else None - 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 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 - a stale asynchronous unwind cannot clear a newer turn. - """ - 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 - 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. - """ - 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 - return True - - def recover_interrupted_turns( - self, - max_age_seconds: int = 60 * 60, - ) -> int: - """Promote crash-left turn markers into ``resume_pending`` (unclean startup only). - - Old/invalid markers are cleared without resuming; suspended sessions - are never re-armed. Returns the number of newly promoted sessions. - """ - 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 - - 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. - 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() - - 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 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.""" - def _apply(entry: SessionEntry): - # Never override an explicit ``suspended`` (hard forced-wipe). - if entry.suspended: - return False - entry.resume_pending = True - entry.resume_reason = reason - entry.last_resume_marked_at = _now() - - return self._update_entry(session_key, _apply) - - def clear_resume_pending(self, session_key: str) -> bool: - """Clear the resume-pending flag after a successful resumed turn. - Returns True if a flag was cleared.""" - def _apply(entry: SessionEntry): - if not entry.resume_pending: - return False - entry.resume_pending = False - entry.resume_reason = None - entry.last_resume_marked_at = None - - return self._update_entry(session_key, _apply) - - def prune_old_entries(self, max_age_days: int) -> int: - """Drop routing entries idle (by ``updated_at``) for more than max_age_days. - - Suspended entries and entries with active background processes are - kept. The SQLite transcript stays; only the key -> session_id mapping - is dropped. ``max_age_days <= 0`` disables. Returns the count removed. - """ - 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 - # 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) - for key in removed_keys: - self._entries.pop(key, None) - if removed_keys: - self._save() - - if removed_keys: - logger.info( - "SessionStore pruned %d entries older than %d days", - len(removed_keys), max_age_days, - ) - return len(removed_keys) - - def suspend_recently_active(self, max_age_seconds: int = 120) -> int: - """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.""" - 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 reset_session(self, session_key: str, display_name: Optional[str] = None) -> Optional[SessionEntry]: """Force reset a session, creating a new session ID.""" @@ -2934,38 +1436,6 @@ class SessionStore: self._save() return new_entry - def advance_compression_session( - self, - session_key: str, - expected_session_id: str, - target_session_id: str, - ) -> Optional[SessionEntry]: - """CAS-advance one route along an already-verified compression lineage. - - Unlike ``switch_session`` this never ends/reopens SQLite rows (the - compression transaction owns that). ``None`` means the route moved - after the caller's snapshot (e.g. /new) — caller must fail closed. - """ - if not session_key or not expected_session_id or not target_session_id: - return None - - with self._lock: - entry = self._entry_locked(session_key) - if entry is None: - return None - if entry.session_id == target_session_id: - 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, - ): - return None - # Bookkeeping, not user activity: leave ``updated_at`` alone. - self._save() - return 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 @@ -3020,10 +1490,7 @@ class SessionStore: return None with self._lock: self._ensure_loaded_locked() - for entry in self._entries.values(): - if entry.session_id == session_id: - return entry - return None + return next((e for e in self._entries.values() if e.session_id == session_id), None) def lookup_by_session_key(self, session_key: str) -> Optional[SessionEntry]: """Return the persisted routing entry for an exact session key.""" @@ -3039,557 +1506,13 @@ class SessionStore: with self._lock: entry = self._entry_locked(session_key) return entry.session_id if entry else None - - def _get_transcript_drain_lock(self): - """Return the lock that serializes pending-queue drain boundaries.""" - return self._lazy("_transcript_drain_lock", threading.RLock) - def append_to_transcript(self, session_id: str, message: Dict[str, Any], skip_db: bool = False) -> None: - """Serialize transcript draining across queue migration boundaries.""" - 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 - ) - def _follow_reroutes(self, session_id: str) -> str: - """Follow the compression reroute chain (cycle-guarded).""" - reroutes = self._lazy("_transcript_reroutes", dict) - seen = set() - while session_id in reroutes and session_id not in seen: - seen.add(session_id) - 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. - """ - 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)) - if spool_path is not None: - self._lazy("_spooled_drop_sessions", set).add(session_id) - logger.warning( - "Session DB transcript pending queue full for %s " - "(cap=%d); spooled oldest message to %s for replay " - "after DB recovery", - session_id, self._MAX_PENDING_PER_SESSION, spool_path, - ) - else: - logger.warning( - "Session DB transcript pending queue full for %s " - "(cap=%d); dropping oldest message to make room " - "(on-disk spool unavailable)", - session_id, self._MAX_PENDING_PER_SESSION, - ) - return pending - - def _divert_transcript_after_db_replaced( - self, session_id: str, queue_session_id: str, exc: Exception - ) -> None: - """Stop SQLite writes on a replaced/quarantined handle and divert the backlog. - - Retrying cannot succeed and the FTS rebuild must never run on this - handle; the pending queue goes to the on-disk spool + JSONL fallback. - """ - logger.error( - "Session DB refused further writes on this handle for " - "%s (%s); stopping SQLite writes and diverting pending " - "transcripts to the on-disk fallback: %s", - session_id, type(exc).__name__, exc, - ) - with self._transcript_retry_lock: - remaining = list(self._dirty_transcripts.get(queue_session_id, [])) - self._dirty_transcripts.pop(queue_session_id, None) - self._transcript_append_failures.pop(session_id, None) - for dropped in remaining: - try: - from gateway.shutdown_flush import spool_dropped_transcript_message - - spool_dropped_transcript_message(session_id, dropped) - except Exception: - logger.warning( - "pending fallback failed for replaced " - "state.db transcript on %s", - session_id, - exc_info=True, - ) - try: - from hermes_state import divert_session_transcript_jsonl - divert_session_transcript_jsonl(session_id, remaining) - except Exception: - logger.warning( - "JSONL divert failed for replaced state.db " - "transcript on %s", - session_id, - exc_info=True, - ) - - def _live_compression_child(self, session_id: str) -> str: - """Transitive compression tip of *session_id* if it is a different, still-live - row, else "" (a depth-1 lookup misses multi-hop lineages). - - Uses the PARENT's proven owner handle: the child's id is not published - until after its write succeeds, so a by-id lookup would fall back to - the ambient store. - """ - owner_db = self._db_for_session_id(session_id) - if owner_db is None: - return "" - tip = owner_db.get_compression_tip(session_id) - if tip and tip != session_id: - tip_row = owner_db.get_session(tip) - if tip_row is not None and tip_row.get("ended_at") is None: - return str(tip) - return "" - - def _migrate_transcript_queue_to_child( - self, session_id: str, queue_session_id: str, child_id: str, pending: list, msg - ) -> list: - """Move the retry queue + failure counter from parent to child and publish - the reroute (retry lock held). Returns the child's pending list. - - Older parent backlog must precede messages already queued directly on - the child. Routing is published only AFTER the queue moved (caller), so - new child writes cannot bypass older parent backlog. - """ - if pending and pending[0] is msg: - pending.pop(0) - existing_child_pending = self._dirty_transcripts.get(child_id, []) - if pending: - pending.extend(existing_child_pending) - self._dirty_transcripts[child_id] = pending - elif existing_child_pending: - pending = existing_child_pending - self._dirty_transcripts.pop(queue_session_id, None) - 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), - ) - self._transcript_reroutes[session_id] = child_id - return pending - - def _publish_transcript_reroute(self, session_id: str, child_id: str) -> None: - """Repoint every route at the compression child and save (index authoritative again).""" - with self._lock: - for entry in self._entries.values(): - if entry.session_id == session_id: - entry.session_id = child_id - self._save() - _hints = getattr(self, "_session_owner_hints", None) - if _hints: - _hints.pop(child_id, None) - - def _append_to_transcript_serialized( - self, session_id: str, message: Dict[str, Any] - ) -> None: - """Append a message to a session's transcript (SQLite), draining the - per-session retry queue.""" - with self._transcript_retry_lock: - pending = self._enqueue_transcript_message(session_id, message) - msg = pending[0] - queue_session_id = session_id - - def _ack_head() -> bool: - """Pop the acknowledged head (retry lock held). True if queue drained.""" - if pending and pending[0] is msg: - pending.pop(0) - if not pending: - self._dirty_transcripts.pop(queue_session_id, None) - self._transcript_append_failures.pop(session_id, None) - return True - return False - - # DB write outside the retry lock so other sessions can append. - while True: - try: - self._append_transcript_message(session_id, msg) - except Exception as exc: - from hermes_state import ( - CompressionSessionClosedError, - StateDbCorruptError, - StateDbReplacedError, - ) - - if isinstance(exc, (StateDbReplacedError, StateDbCorruptError)): - self._divert_transcript_after_db_replaced(session_id, queue_session_id, exc) - return - - if isinstance(exc, CompressionSessionClosedError): - # Adopt only a different, still-live compression tip, else - # fail closed. - _owner_key = self._owner_key_for_session_id(session_id) - child_id = self._live_compression_child(session_id) - if child_id: - # Record the child's owner BEFORE writing to it (the - # reroute is published only after the write succeeds - # — load-bearing for backlog order). - if _owner_key: - self._lazy("_session_owner_hints", dict)[child_id] = _owner_key - try: - self._append_transcript_message(child_id, msg) - except Exception as reroute_exc: - exc = reroute_exc - else: - with self._transcript_retry_lock: - pending = self._migrate_transcript_queue_to_child( - session_id, queue_session_id, child_id, pending, msg - ) - queue_session_id = child_id - self._publish_transcript_reroute(session_id, child_id) - if not pending: - return - msg = pending[0] - session_id = child_id - continue - else: - # Permanent routing invariant failure, not a transient - # outage: drop it so it cannot poison later writes. - with self._transcript_retry_lock: - _ack_head() - logger.error( - "Session DB transcript append rejected for compression-ended " - "%s with no unique live child; not retrying", - session_id, - ) - return - if self._is_fts_corruption_error(exc) and self._rebuild_fts_once(): - try: - self._append_transcript_message(session_id, msg) - except Exception as retry_exc: - exc = retry_exc - else: - with self._transcript_retry_lock: - _ack_head() - continue - with self._transcript_retry_lock: - failures = self._transcript_append_failures.get(session_id, 0) + 1 - self._transcript_append_failures[session_id] = failures - logger.warning( - "Session DB transcript append failed for %s " - "(failure_count=%d, pending=%d); will retry: %s", - session_id, failures, len(pending), exc, - ) - return - else: - with self._transcript_retry_lock: - queue_empty = _ack_head() - if not queue_empty: - msg = pending[0] - if queue_empty: - # Backlog clear: replay cap-dropped messages spooled to disk. - self._drain_spooled_drops(session_id) - return - continue - - def _drain_spooled_drops(self, session_id: str) -> None: - """Replay cap-dropped spooled transcript messages after DB recovery. - - Best-effort: replay failures keep the spool files for the next - successful flush; nothing here may raise into the caller. - """ - spooled_sessions = getattr(self, "_spooled_drop_sessions", None) - if not spooled_sessions or session_id not in spooled_sessions: - return - try: - from gateway.shutdown_flush import drain_transcript_spool - - _replayed, remaining = drain_transcript_spool( - session_id, - lambda message: self._append_transcript_message( - session_id, message - ), - ) - if not remaining: - spooled_sessions.discard(session_id) - except Exception as exc: - logger.warning("Failed to drain transcript spool for %s: %s", session_id, exc) - - def _append_transcript_message(self, session_id: str, message: Dict[str, Any]) -> None: - """Write one transcript row. Caller handles retry queuing.""" - _db = self._db_for_session_id(session_id) - if _db is None: - # Named profile with no resolvable home yet: defer (caller queues) - # instead of writing into the ambient store. - raise RuntimeError( - f"no owning session store for {session_id}; deferring transcript write" - ) - is_assistant = message.get("role") == "assistant" - _db.append_message( - session_id=session_id, - role=message.get("role", "unknown"), - content=message.get("content"), - tool_name=message.get("tool_name"), - tool_calls=message.get("tool_calls"), - tool_call_id=message.get("tool_call_id"), - reasoning=message.get("reasoning") if is_assistant else None, - reasoning_content=message.get("reasoning_content") if is_assistant else None, - reasoning_details=message.get("reasoning_details") if is_assistant else None, - codex_reasoning_items=message.get("codex_reasoning_items") if is_assistant else None, - codex_message_items=message.get("codex_message_items") if is_assistant else None, - platform_message_id=(message.get("platform_message_id") or message.get("message_id")), - observed=bool(message.get("observed")), - timestamp=message.get("timestamp"), - # Exact bytes sent to the API (prompt-cache-stable replay); must - # survive every persistence path or the next replay diverges. - api_content=extract_api_content_sidecar(message), - # Presentation typing (e.g. "internal_notification"); DB-only. - display_kind=message.get("display_kind"), - display_metadata=message.get("display_metadata"), - ) # Max in-memory pending messages per session (DB persistently broken). - _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. - A bare SQLITE_CORRUPT can mean structural B-tree damage; only errors - naming ``messages_fts`` or carrying FTS provenance (per - ``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: - return True - import sqlite3 - from hermes_state import SessionDB - - if isinstance(exc, sqlite3.DatabaseError): - return SessionDB._is_fts_write_corruption_error(exc) - return False - - def _rebuild_fts_once(self) -> bool: - """Attempt FTS5 ``rebuild`` once per store lifetime; True if any index was rebuilt.""" - if self._fts_rebuild_attempted: - return False - self._fts_rebuild_attempted = True - db = self._db - if db is None or not hasattr(db, "rebuild_fts"): - return False - # WAL split-brain guard: skip when a foreign process holds state.db. - if hasattr(db, "_foreign_state_db_holders"): - foreign_holders = db._foreign_state_db_holders() - if foreign_holders: - logger.warning( - "Skipping Session DB FTS rebuild while foreign processes " - "hold the database or WAL sidecars (%s); canonical " - "transcript writes remain available.", - foreign_holders, - ) - return False - try: - rebuilt = db.rebuild_fts() - except Exception as exc: - 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, - ) - return rebuilt > 0 - - def _clear_dirty_transcript(self, session_id: str) -> None: - """Drop queued pending messages so a rewrite/rewind doesn't re-insert them.""" - with self._transcript_retry_lock: - self._dirty_transcripts.pop(session_id, None) - self._transcript_append_failures.pop(session_id, None) - - def has_platform_message_id( - self, session_id: str, platform_message_id: str - ) -> bool: - """Whether a message with this platform_message_id is persisted (False without a DB).""" - db = self._db_for_session_id(session_id) - if not db: - return False - try: - 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 - - def rewrite_transcript( - self, - session_id: str, - messages: List[Dict[str, Any]], - active_only: bool = False, - reject_active_turn_lease: bool = False, - ) -> bool: - """Replace a session's transcript (/retry, /compress). - - DESTRUCTIVE by default: ``active_only=False`` DELETEs every row - including soft-archived compaction history; pass ``active_only=True`` - for sessions that may carry archived rows. Returns ``True`` when the - write lands (or there is no DB), ``False`` on failure — callers about - to commit a destructive change on top (e.g. /compress repointing) - must check it. ``reject_active_turn_lease`` is for user-initiated - rewrites that do not own the cross-process turn lease. - """ - db = self._db_for_session_id(session_id) - if not db: - return True - with self._get_transcript_drain_lock(): - try: - db.replace_messages( - session_id, - messages, - active_only=active_only, - reject_active_turn_lease=reject_active_turn_lease, - ) - except Exception as e: - logger.debug("Failed to rewrite transcript in DB: %s", e) - return False - self._clear_dirty_transcript(session_id) - return True - - def load_transcript(self, session_id: str) -> List[Dict[str, Any]]: - """Load all messages from a session's transcript (state.db is canonical). - - Reads follow the same routing writes use: the in-memory reroute map - (compression rotation), then the durable compression tip — otherwise - the transcript "vanishes" while every message sits under the child. - """ - if not self._db_for_session_id(session_id): - return [] - session_id = self._follow_reroutes(session_id) - try: - # Durable successor survives restart; the reroute map doesn't. - tip = self._db_for_session_id(session_id).get_compression_tip(session_id) - if tip: - session_id = tip - except Exception: - pass - try: - # repair_alternation: this feeds LIVE REPLAY; heal a durable - # user;user wedge once here instead of on every request. - return self._db_for_session_id(session_id).get_messages_as_conversation( - session_id, repair_alternation=True - ) - except Exception as e: - # Empty history is valid data; a failed canonical read is not — - # live-replay callers must fail closed, not start from []. - logger.error( - "Transcript read failed for session %s; refusing to treat the " - "conversation as empty: %s", - session_id, - e, - exc_info=True, - ) - raise TranscriptReadError(session_id) from e - - def rewind_session( - self, - session_id: str, - n: int = 1, - *, - require_retryable_composite: bool = False, - ) -> Optional[Dict[str, Any]]: - """Back up ``n`` user turns via soft-delete (``active=0``), mirroring CLI ``/undo [N]``. - - Returns ``{"rewound_count", "turns_undone", "target_text"}`` or ``None`` - (no DB / no user turn). ``n`` clamps to the oldest user turn. - ``require_retryable_composite`` is the gateway ``/retry`` guard: the - selected turn must be a composite carrier whose live payload is - losslessly replayable as text before anything changes. - """ - db = self._db_for_session_id(session_id) - if not db: - return None - with self._get_transcript_drain_lock(): - if n < 1: - n = 1 - from agent.context_compressor import ( - retryable_user_text, - split_user_originated_turn, - user_originated_turn_view, - ) - - try: - expected_active_ids = db.get_active_message_ids(session_id) - durable = db.get_messages_as_conversation( - session_id, - include_row_ids=True, - ) - user_indices = [ - index - for index, message in enumerate(durable) - if user_originated_turn_view(message) is not None - ] - if not user_indices: - return None - turns_undone = min(n, len(user_indices)) - target = durable[user_indices[-turns_undone]] - target_id = target.get("_row_id") - if not isinstance(target_id, int): - return None - handoff, target_view = split_user_originated_turn(target) - if target_view is None: - return None - if require_retryable_composite and handoff is None: - return None - except Exception as e: - logger.debug("rewind_session: failed to resolve canonical target: %s", e) - return None - if require_retryable_composite: - # Keep replay-policy failures distinct from persistence errors - # so /retry can explain why the selected carrier is unsafe. - target_text = retryable_user_text(target_view.get("content")) - try: - result = db.rewind_to_message( - session_id, - target_id, - preserve_compaction_handoff=handoff is not None, - expected_active_ids=expected_active_ids, - expected_target_content=target_view.get("content"), - ) - except ValueError as e: - logger.debug("rewind_session: %s", e) - return None - except Exception as e: - logger.debug("rewind_session: rewind_to_message failed: %s", e) - return None - self._clear_dirty_transcript(session_id) - # ``target_view`` is the live projection; a composite carrier's raw - # row holds the summary wrapper and must not be echoed as prompt. - if not require_retryable_composite: - content = target_view.get("content") or "" - if isinstance(content, list): - parts = [ - p.get("text", "") - for p in content - if isinstance(p, dict) and p.get("type") == "text" - ] - target_text = "\n".join(t for t in parts if t) - elif isinstance(content, str): - target_text = content - else: - target_text = "" - return { - "rewound_count": result.get("rewound_count", 0), - "turns_undone": turns_undone, - "target_text": target_text, - } def build_session_context( @@ -3599,13 +1522,11 @@ 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 - context = SessionContext( source=source, connected_platforms=connected, @@ -3616,11 +1537,9 @@ def build_session_context( thread_sessions_per_user=getattr(config, "thread_sessions_per_user", False), ), ) - if session_entry: context.session_key = session_entry.session_key context.session_id = session_entry.session_id context.created_at = session_entry.created_at context.updated_at = session_entry.updated_at - return context diff --git a/gateway/session_lifecycle.py b/gateway/session_lifecycle.py new file mode 100644 index 0000000000..4497cce47d --- /dev/null +++ b/gateway/session_lifecycle.py @@ -0,0 +1,385 @@ +"""SessionStore reset/expiry policy and crash-recovery markers: idle/daily reset +evaluation, expiry finalization, active-turn tokens, resume_pending, +suspension and pruning. + +Mixin split out of ``gateway/session.py``; bound onto ``SessionStore`` via the MRO. +""" + +from __future__ import annotations + +import logging +import uuid +from datetime import datetime, timedelta +from typing import TYPE_CHECKING, Optional + +if TYPE_CHECKING: + from gateway.session import SessionEntry, SessionSource + +# Log-record parity with the origin module. +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 set_expiry_finalized( + self, entry: SessionEntry, *, clear_model_override: bool = True + ) -> None: + """Mark a session entry expiry-finalized in memory, sessions.json, AND state.db. + + Single write-path for the expiry watcher so the durable flag survives + sessions.json loss. ``clear_model_override=False`` = flag only. + """ + with self._lock: + entry.expiry_finalized = True + if clear_model_override: + # Finalization is a conversation boundary: drop the persisted + # /model override so a later message cannot rehydrate it. + entry.model_override = None + self._save() + # 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) + 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, + ) + if now.hour < policy.at_hour: + today_reset -= timedelta(days=1) + if updated_at < today_reset: + return "daily" + return None + + def _is_session_expired(self, entry: SessionEntry) -> bool: + """Whether the entry's reset policy has expired it (entry alone, no source). + + Used by the background expiry watcher. Sessions with active + background processes are never considered expired. + """ + 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, + ) + return self._policy_reset_reason(policy, entry.updated_at) is not None + + def is_session_finalizable(self, entry: SessionEntry) -> bool: + """True if the expiry watcher will *ever* finalize this session. + + A ``mode == "none"`` session never expires, so the agent-cache idle + sweep must reap its agent itself instead of deferring to the watcher + (deferring would pin the agent for the gateway's lifetime). Policy + resolution errors count as "not finalizable" (sweep reaps — safe). + """ + try: + policy = self.config.get_reset_policy( + platform=entry.platform, + session_type=entry.chat_type, + ) + return policy.mode != "none" + except Exception: + return False + + def _is_session_ended_in_db(self, session_id: str) -> bool: + """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 + session ended while the gateway stays alive. Store resolved from the + row's owning profile, not the ambient scope. + """ + db = self._db_for_session_id(session_id) + if not db or not session_id: + return False + try: + row = db.get_session(session_id) + except Exception: + return False + 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. + """ + 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 + ) + return self._policy_reset_reason(policy, entry.updated_at) + + def _route_reset_reason( + self, entry: SessionEntry, source: SessionSource, now: datetime + ) -> Optional[str]: + """Reset decision for an existing route (no lock; DB/config I/O). + + ``suspended`` always resets. Otherwise the reset policy decides; a + still-pending resume marker is additionally freshness-gated — but + ``session_reset.mode: none`` (user opted out of ALL automatic resets) + 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, + ) + if policy.mode == "none": + return None + window = auto_continue_freshness_window() + ref_time = entry.last_resume_marked_at or entry.updated_at + if window > 0 and (now - ref_time).total_seconds() > window: + return "resume_pending_expired" + return None + + def _update_entry(self, session_key: str, mutate) -> bool: + """Apply ``mutate(entry)`` under ``_lock`` and full-save; False when the + entry is missing or *mutate* returned False (nothing to persist).""" + with self._lock: + entry = self._entry_locked(session_key) + if entry is None or mutate(entry) is False: + return False + self._save() + return True + + 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 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 + 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 + 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. + """ + 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 + return True + + def recover_interrupted_turns( + self, + max_age_seconds: int = 60 * 60, + ) -> int: + """Promote crash-left turn markers into ``resume_pending`` (unclean startup only). + + 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 + + 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. + 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() + + 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 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: + return False + entry.resume_pending = True + entry.resume_reason = reason + entry.last_resume_marked_at = _now() + + return self._update_entry(session_key, _apply) + + def clear_resume_pending(self, session_key: str) -> bool: + """Clear the resume-pending flag after a successful resumed turn. + Returns True if a flag was cleared.""" + def _apply(entry: SessionEntry): + if not entry.resume_pending: + return False + entry.resume_pending = False + entry.resume_reason = None + entry.last_resume_marked_at = None + + return self._update_entry(session_key, _apply) + + def prune_old_entries(self, max_age_days: int) -> int: + """Drop routing entries idle (by ``updated_at``) for more than max_age_days. + + Suspended entries and entries with active background processes are + 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 + # 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) + for key in removed_keys: + self._entries.pop(key, None) + if removed_keys: + self._save() + + if removed_keys: + logger.info( + "SessionStore pruned %d entries older than %d days", + len(removed_keys), max_age_days, + ) + return len(removed_keys) + + def suspend_recently_active(self, max_age_seconds: int = 120) -> int: + """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 diff --git a/gateway/session_persistence.py b/gateway/session_persistence.py new file mode 100644 index 0000000000..9d1f044da1 --- /dev/null +++ b/gateway/session_persistence.py @@ -0,0 +1,656 @@ +"""SessionStore storage plumbing: per-profile SessionDB handle resolution and the +routing-index load/save paths (state.db gateway_routing primary, sessions.json +legacy mirror). + +Mixin split out of ``gateway/session.py``; bound onto ``SessionStore`` via the MRO. +""" + +from __future__ import annotations + +import logging +import json +import os +import threading +from pathlib import Path +from typing import Any, Dict, Optional +from utils import atomic_replace +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from gateway.session import SessionEntry + +# Log-record parity with the origin module. +logger = logging.getLogger("gateway.session") + + +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). + """ + + def _open_session_db_for_active_scope(self, db_path: Optional[Path] = None): + """SessionDB for the profile scope active on this task. + + ``db_path`` pins the store; otherwise ``_default_db_path()`` follows the + context-local HERMES_HOME from ``_profile_runtime_scope`` (resolved per + call so multiplexed profiles reach their own store). Handles are cached + per path; failed opens enter a bounded backoff during which callers keep + using the JSONL fallback. + """ + 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}") + raise + + return self._db_handle_cache.get( + path, + _open, + non_cacheable=lambda exc: ( + isinstance(exc, RuntimeError) and "live-system guard" in str(exc) + ), + ) + + 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 + def _db(self): + """The SessionDB for the active profile scope, or a pinned override. + + Assigning ``store._db`` pins that value for every subsequent read + (tests install a fake or disable the DB with ``store._db = None``). + 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 + return self._open_session_db_for_active_scope() + + @_db.setter + def _db(self, value) -> None: + self._db_pinned = value + + @property + def _routing_db(self): + """The one store that owns the routing index, whatever scope is active. + + ``_entries`` is one flat dict holding every profile's keys, so it must + persist to ONE file (``_routing_home``), not whichever profile is + scoped — otherwise a mid-turn rewrite and the unscoped startup load see + different copies and crash markers under a secondary profile go + 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 + home = getattr(self, "_routing_home", None) + try: + if home is None: + return self._db + return self._open_session_db_for_active_scope(db_path=home / "state.db") + except Exception: + return None + + def _named_profile_for_key(self, session_key: Optional[str]) -> Optional[str]: + """The non-default profile that owns *session_key*, or None. + + None means the ambient store is authoritative (multiplexing off, or + legacy ``agent:main`` namespace). It deliberately does NOT cover "that + profile has no directory" — ownership and resolvability are separate + questions that ``_db_for_key`` answers separately. + """ + if not getattr(self.config, "multiplex_profiles", False): + return None + profile = self._profile_from_session_key(session_key) + if not profile or profile == "default": + return None + 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. + """ + profile = self._named_profile_for_key(session_key) + if profile is None: + return None + cache = self._profile_home_cache + if profile in cache: + return cache[profile] + home: Optional[Path] = None + try: + from hermes_cli.profiles import get_profile_dir, profile_exists + + if profile_exists(profile): + home = Path(get_profile_dir(profile)) + except Exception as exc: + logger.debug("Could not resolve profile home for %r: %s", session_key, exc) + home = None + # Only hits are memoized: a profile directory can be provisioned + # *after* startup (enrollment bridge), and a cached miss would pin + # that profile's rows to the ambient store for the process lifetime. + if home is not None: + cache[profile] = home + return home + + def _db_for_key(self, session_key: Optional[str]): + """The SessionDB holding *session_key*'s rows, whatever scope is active. + + ``_db`` follows the ambient HERMES_HOME that only the inbound message + path installs; background work (expiry watcher) runs unscoped and would + 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 + profile = self._named_profile_for_key(session_key) + if profile is None: + return self._db + home = self._profile_home_for_key(session_key) + if home is None: + # Named owner we cannot resolve (not provisioned yet, or lookup + # failed). Falling back to the ambient store would split ONE + # session identity across two physical stores — fail closed; + # callers already handle a missing DB. + logger.warning( + "gateway.session: profile %r has no resolvable home (key %r); " + "refusing to fall back to the ambient store", + profile, session_key, + ) + return None + try: + return self._open_session_db_for_active_scope(db_path=home / "state.db") + except Exception: + # Same contract as ``_db``: a failed open degrades to JSONL fallback. + return None + + def _owner_key_for_session_id(self, session_id: Optional[str]) -> Optional[str]: + """The routing key that owns *session_id*, or None. + + The published index is authoritative; ``_session_owner_hints`` covers + the window where ownership is proven but routing not yet published. + Deliberately lock-free: several callers already hold ``_lock``. + """ + if not session_id: + return None + try: + for entry in list(self._entries.values()): + if entry.session_id == session_id: + return entry.session_key + except Exception: + pass # bare stores / foreign entry objects in suites + return (getattr(self, "_session_owner_hints", None) or {}).get(session_id) + + def _db_for_session_id(self, session_id: Optional[str]): + """The SessionDB holding *session_id*'s row (owner recovered from the + index or a pre-published hint; unknown ids fall back to the ambient store).""" + if not session_id: + return self._db + return self._db_for_key(self._owner_key_for_session_id(session_id)) + + def close_all_db_handles(self) -> None: + """Close every SessionDB handle this store opened (one per path). + + Closing only ``store._db`` would strand secondary profiles' handles with + their WAL lock held ('database is locked' on restart). Drained under the + lock, closed outside it; a pinned handle is the pinner's to close. + """ + def _close(db) -> None: + # Shared instances no-op on close(); release the refcount instead. + from hermes_state import release_or_close + try: + release_or_close(db) + except Exception as exc: + logger.debug("SessionDB close error during handle sweep: %s", exc) + + self._db_handle_cache.close_all(_close) + + def _ensure_loaded(self) -> None: + """Load sessions index from disk if not already loaded.""" + with self._lock: + self._ensure_loaded_locked() + + def _entry_locked(self, session_key: str) -> Optional[SessionEntry]: + """Load the index and return the entry for *session_key*. Lock held.""" + self._ensure_loaded_locked() + return self._entries.get(session_key) + + def _routing_scope(self) -> str: + """Namespace for this store's gateway_routing rows: the resolved + sessions_dir, so stores with different dirs never share entries.""" + try: + return str(Path(self.sessions_dir).resolve()) + except Exception: + return str(self.sessions_dir) + + def _routing_db_method(self, name: str): + """Bound ``_routing_db.`` if the handle exists and has it, else None.""" + db = self._routing_db + method = getattr(db, name, None) if db else None + return method if callable(method) else None + + @staticmethod + def _routing_entry_from_json(key: str, entry_json: str) -> Optional[SessionEntry]: + """Parse one gateway_routing row; None (with a warning) when invalid.""" + from gateway.session import SessionEntry + try: + entry_data = json.loads(entry_json) + if isinstance(entry_data, dict): + return SessionEntry.from_dict(entry_data) + except (ValueError, KeyError, TypeError) as e: + logger.warning("Skipping invalid routing entry %r: %s", key, e) + return None + + def _ensure_loaded_locked(self) -> None: + """Load the routing index. Must be called with self._lock held. + + state.db ``gateway_routing`` is primary; sessions.json is the legacy + import for keys the DB lacks (persisted to the DB on the next _save). + """ + if self._loaded: + self._reconcile_recovered_routing_locked() + return + + 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) + + 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()} + ) + + # A hard crash skips graceful shutdown and leaves sessions.json + # pointing at ended sessions; self-heal before the first message. + self._prune_stale_sessions_locked() + + def _import_legacy_sessions_json(self, db_had_entries: bool) -> None: + """Legacy import: sessions.json fills only keys the DB lacks. Lock held.""" + from gateway.session import SessionEntry + sessions_file = self.sessions_dir / "sessions.json" + if not sessions_file.exists(): + return + try: + with open(sessions_file, "r", encoding="utf-8") as f: + data = json.load(f) + imported = 0 + for key, entry_data in data.items(): + # "_"-prefixed keys are sentinels (e.g. "_README"), not entries. + if key.startswith("_") or key in self._entries: + continue + # A non-dict entry (corrupt file) must not abort the whole load. + if not isinstance(entry_data, dict): + logger.warning( + "Skipping invalid session entry %r: " + "expected dict, got %s", + key, type(entry_data).__name__, + ) + continue + try: + self._entries[key] = SessionEntry.from_dict(entry_data) + imported += 1 + except (ValueError, KeyError, TypeError) as e: + logger.warning("Skipping invalid session entry %r: %s", key, e) + if imported and db_had_entries: + logger.info( + "gateway.session: imported %d legacy sessions.json " + "entr%s missing from state.db routing table", + imported, "y" if imported == 1 else "ies", + ) + except Exception as e: + print(f"[gateway] Warning: Failed to load sessions: {e}") + + def _prune_stale_sessions_locked(self) -> None: + """Remove routing entries whose session has ended in state.db (startup, lock held). + + Stale == ``end_reason IS NOT NULL``. Rows absent from the DB are kept; + a ``None`` DB handle is a no-op; DB errors are non-fatal. + """ + if not self._entries: + return + + stale_keys: list = [] + recovered_keys = 0 + try: + for key, entry in self._entries.items(): + # Ask the store that owns the key, not the ambient handle, or a + # live secondary-profile session gets pruned on the root copy. + db = self._db_for_key(key) + if db is None: + continue + row = db.get_session(entry.session_id) + if row is None or row.get("end_reason") is None: + continue + verdict = self._stale_entry_verdict(key, entry, row) + if verdict == "prune": + stale_keys.append(key) + elif verdict is not None: + self._entries[key] = verdict + recovered_keys += 1 + except Exception as exc: + logger.warning( + "gateway.session: stale-entry pruning skipped due to DB error: %s", + exc, + ) + return + + for key in stale_keys: + del self._entries[key] + + if stale_keys or recovered_keys: + self._save() + + def _stale_entry_verdict(self, key: str, entry, row): + """For a routing entry whose row has ended: ``"prune"``, a replacement + entry (repoint), or None (keep as-is).""" + from gateway.session import _now + recovered_entry = None + if entry.origin is not None: + try: + recovered_entry = self._recover_session_from_db( + session_key=key, + source=entry.origin, + now=_now(), + raise_on_lookup_error=True, + ) + except Exception as exc: + # Indeterminate: keep the only routing handle. + logger.debug( + "gateway.session: recovery lookup failed for stale " + "sessions.json entry %r -> %s: %s", + key, + entry.session_id, + exc, + ) + return None + + # Compression-ended parent with a newer live child for the same peer: + # repoint instead of dropping, or queued/resume-pending work vanishes + # until the next message. + if recovered_entry is not None and recovered_entry.session_id != entry.session_id: + 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, + ) + return recovered_entry + + # Same-id recovery == successful resume: keep the ORIGINAL entry object + # (the recovered one is rebuilt minimal and would drop counters, + # model_override, resume markers, metadata). Nothing changes, no save. + if recovered_entry is not None: + logger.info( + "gateway.session: reopened ended session %s for " + "sessions.json entry %r (end_reason=%r); keeping route", + entry.session_id, key, row["end_reason"], + ) + return None + + logger.warning( + "gateway.session: pruning stale sessions.json entry " + "%r -> %s (end_reason=%r); left by a crashed gateway", + key, entry.session_id, row["end_reason"], + ) + return "prune" + + def _save(self) -> None: + """Persist the routing index while the caller holds ``_lock``.""" + data, generation = self._snapshot_routing_locked() + self._persist_routing_data(data, generation) + + def _next_routing_generation_locked(self) -> int: + """Bump and return the shared routing counter. Caller holds ``_lock``. + + Full snapshots AND single-entry fast saves MUST allocate from this one + counter: the stale-write protection is a total order over + serialization times and silently breaks otherwise. + """ + self._routing_generation = getattr(self, "_routing_generation", 0) + 1 + return self._routing_generation + + def _reconcile_recovered_routing_locked(self) -> None: + """Merge authoritative rows after a fallback-only startup load.""" + baseline = getattr(self, "_routing_fallback_baseline", None) + if getattr(self, "_routing_db_loaded", False) or baseline is None: + return + + loader = self._routing_db_method("load_gateway_routing_entries") + if loader is None: + return + try: + durable = loader(scope=self._routing_scope()) + except Exception as exc: + 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()} + for key, entry_json in durable.items(): + durable_entry = self._routing_entry_from_json(key, entry_json) + if durable_entry is None: + continue + + if key not in baseline: + # A key created while on fallback wins over a DB-only key; + # otherwise restore the authoritative row that fallback never saw. + self._entries.setdefault(key, durable_entry) + elif key not in current: + # The key was loaded from fallback and deliberately removed. + continue + elif current[key] == baseline[key]: + # Unchanged fallback data yields to the authoritative DB copy. + self._entries[key] = durable_entry + + self._routing_db_loaded = True + self._routing_fallback_baseline = None + + 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(), + ) + + def _persist_routing_data(self, data: Dict[str, Any], generation: int) -> None: + """Serialize all whole-index writers through one durable write lock.""" + with self._lazy("_save_lock", threading.Lock): + if generation <= getattr(self, "_persisted_routing_generation", 0): + return + # Fold in fast upserts numbered above this snapshot: they were + # serialized after us and a delayed full rewrite must not regress them. + fast_persisted = getattr(self, "_fast_persisted_entries", None) + if fast_persisted: + for key, (revision, entry_json) in fast_persisted.items(): + if revision > generation: + data[key] = json.loads(entry_json) + db_saved = False + 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(), + ) + db_saved = True + except Exception as exc: + logger.warning("gateway.session: state.db routing save failed: %s", exc) + if getattr(self, "_write_sessions_json", True) or not db_saved: + try: + self._save_sessions_json(data) + except Exception as exc: + if not db_saved: + raise + # state.db is authoritative. A failed legacy mirror must not + # report the already-committed primary write as failed. + logger.warning( + "gateway.session: sessions.json mirror save failed " + "after state.db commit: %s", + exc, + ) + self._persisted_routing_generation = generation + # 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 + ]: + 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 + 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_" + ) + try: + with os.fdopen(fd, "w", encoding="utf-8") as f: + json.dump(data, f, indent=2) + f.flush() + os.fsync(f.fileno()) + atomic_replace(tmp_path, sessions_file) + except BaseException: + try: + os.unlink(tmp_path) + except OSError as e: + logger.debug("Could not remove temp file %s: %s", tmp_path, e) + raise + + def _save_entries(self) -> None: + """Snapshot latest state under ``_lock`` and persist after releasing it.""" + with self._lock: + data, generation = self._snapshot_routing_locked() + self._persist_routing_data(data, generation) + + def _save_entry( + self, + session_key: str, + *, + entry_data: Optional[Dict[str, Any]] = None, + lock_held: bool = False, + ) -> None: + """Persist ONE routing entry via UPSERT — the per-turn fast path + (a full rewrite fsyncs a multi-MB sessions.json, ~50ms at ~1100 keys). + + Invariants: the key -> session_id mapping never changes here — + structural transitions (create/recover/reset/switch/prune/heal) use the + full rewrite, which also refreshes the sessions.json mirror (it may lag + in metadata only; state.db stays primary). The entry is serialized under + ``_lock`` with a revision from the shared routing generation counter + (higher == same-or-newer); under ``_save_lock`` the upsert is skipped if + a full snapshot or a newer fast save of this key already persisted (the + reverse case lives in ``_persist_routing_data``). No DB or a failed + upsert falls back to the full rewrite so DB-less installs stay durable. + ``entry_data`` persists a candidate BEFORE it is published to the live + entry (failure-atomic transitions); the fallback carries the same candidate. + """ + def _capture() -> Optional[tuple[str, int, Optional[Dict[str, Any]]]]: + 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() + ) + 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: + with self._lock: + captured = _capture() + if captured is None: + return + entry_json, revision, candidate_entry = captured + saver = self._routing_db_method("save_gateway_routing_entry") + if saver is not None: + try: + with self._lazy("_save_lock", threading.Lock): + if getattr(self, "_persisted_routing_generation", 0) >= revision: + return + fast_persisted = self._lazy("_fast_persisted_entries", dict) + persisted = fast_persisted.get(session_key) + if persisted is not None and persisted[0] >= revision: + return + saver(session_key, entry_json, scope=self._routing_scope()) + fast_persisted[session_key] = (revision, entry_json) + return + except Exception as exc: + logger.warning( + "gateway.session: single-entry routing save failed for %r " + "(%s); falling back to full index rewrite", + session_key, exc, + ) + 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[session_key] = candidate_entry + self._persist_routing_data(fallback_data, revision) + else: + self._save_entries() diff --git a/gateway/session_recovery.py b/gateway/session_recovery.py new file mode 100644 index 0000000000..892c57b483 --- /dev/null +++ b/gateway/session_recovery.py @@ -0,0 +1,519 @@ +"""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). + +Mixin split out of ``gateway/session.py``; bound onto ``SessionStore`` via the MRO. +""" + +from __future__ import annotations + +import logging +import json +import threading +from dataclasses import replace +from datetime import datetime +from gateway.config import Platform +from typing import TYPE_CHECKING, Any, Dict, Optional + +if TYPE_CHECKING: + from gateway.session import SessionEntry, SessionSource + +# Log-record parity with the origin module. +logger = logging.getLogger("gateway.session") + + +def _origin_json(source) -> Optional[str]: + """``source.to_dict()`` as JSON, or None when absent/unserializable.""" + if source is None: + return None + try: + return json.dumps(source.to_dict()) + except Exception: + return None + + +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). + """ + + def _resolve_profile_for_key(self, source: Optional[SessionSource] = None) -> Optional[str]: + """Profile namespace for session keys: None when multiplexing is off + (legacy ``agent:main``), else ``source.profile`` or the active profile.""" + if not getattr(self.config, "multiplex_profiles", False): + return None + if source is not None and source.profile: + return source.profile + try: + from hermes_cli.profiles import get_active_profile_name + return get_active_profile_name() or "default" + except Exception: + return None + + @staticmethod + def _profile_from_session_key(session_key: Optional[str]) -> Optional[str]: + """Extract the profile namespace encoded in a gateway session key.""" + if not session_key: + return None + parts = str(session_key).split(":") + if len(parts) < 2 or parts[0] != "agent": + return None + namespace = parts[1] or "main" + return "default" if namespace == "main" else namespace + + @staticmethod + def _active_profile_name() -> str: + try: + from hermes_cli.profiles import get_active_profile_name + return get_active_profile_name() or "default" + except Exception: + return "default" + + def _recovered_row_allowed_for_active_profile( + self, + *, + requested_session_key: str, + recovered: Dict[str, Any], + ) -> bool: + """Prevent a gateway from reviving another profile's row. + + Single-profile: the row's namespace must match the ACTIVE profile. + Multiplexed: it must match the namespace of the requested key (the + active profile is meaningless there). Keyless rows stay adoptable. + """ + recovered_key = str(recovered.get("session_key") or "") + if not recovered_key or recovered_key == requested_session_key: + return True + + recovered_profile = self._profile_from_session_key(recovered_key) + if recovered_profile is None: + return True + + if getattr(self.config, "multiplex_profiles", False): + requested_profile = self._profile_from_session_key(requested_session_key) + return requested_profile is None or recovered_profile == requested_profile + + return recovered_profile == self._active_profile_name() + + def _generate_session_key(self, source: SessionSource, key_source: Optional[SessionSource] = None) -> str: + """Session key for *source* (profile resolved from *source*, key built + from *key_source* when given).""" + from gateway.session import build_session_key + return build_session_key( + key_source if key_source is not None else source, + group_sessions_per_user=getattr(self.config, "group_sessions_per_user", True), + thread_sessions_per_user=getattr(self.config, "thread_sessions_per_user", False), + profile=self._resolve_profile_for_key(source), + ) + + def _legacy_slack_session_key(self, source: SessionSource) -> Optional[str]: + """Pre-workspace Slack key for an explicitly scoped source. + + Deliberately Slack-only; an unscoped Slack session may be claimed by + only one workspace because its old key cannot distinguish teams. + """ + 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) + ) + + def _claim_legacy_slack_key(self, legacy_key: Optional[str]) -> bool: + """Atomically reserve one ambiguous legacy Slack key for migration.""" + if not legacy_key: + return False + with self._lazy("_legacy_slack_claim_lock", threading.Lock): + claimed = self._lazy("_claimed_legacy_slack_keys", set) + if legacy_key in claimed: + return False + claimed.add(legacy_key) + return True + + @staticmethod + def _recovered_row_matches_source_scope( + recovered: Dict[str, Any], source: SessionSource + ) -> bool: + """Reject recovered rows whose recorded origin belongs to another workspace. + + A workspace-scoped Slack lookup adopts a row only if its origin_json + 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 + ): + return True + try: + origin = json.loads(recovered.get("origin_json") or "") + except (TypeError, ValueError): + return False + if not isinstance(origin, dict): + return False + return origin.get("scope_id", origin.get("guild_id")) == source.scope_id + + def _create_entry_from_recovered_row( + self, + *, + row: Dict[str, Any], + session_key: str, + source: SessionSource, + now: datetime, + ) -> SessionEntry: + from gateway.session import SessionEntry + def _ts(value, default: datetime) -> datetime: + try: + return datetime.fromtimestamp(float(value)) + except (TypeError, ValueError, OSError): + return default + + # An invalid durable timestamp must look old, never freshly active. + created_at = _ts(row.get("started_at"), datetime.fromtimestamp(0)) + # The finder already returns durable recency; no extra round-trip. + last_activity = row.get("last_activity_at") + 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 + ) + return SessionEntry( + session_key=session_key, + session_id=str(row["id"]), + created_at=created_at, + updated_at=updated_at, + origin=source, + display_name=source.chat_name, + platform=source.platform, + chat_type=source.chat_type, + reset_had_activity=bool(had_activity), + ) + + def _find_gateway_session_row( + self, + *, + session_key: str, + source: SessionSource, + allow_peer_fallback: bool, + raise_on_lookup_error: bool = False, + ) -> Optional[Dict[str, Any]]: + """Query one durable gateway session row. + + Scoped Slack lookups disable SessionDB's platform/chat/user fallback: + that tuple does not contain a workspace id and could therefore revive + another team's session. The caller performs one explicit exact lookup + of the old unscoped key instead. + """ + db = self._db_for_key(session_key) + finder = getattr(db, "find_latest_gateway_session_for_peer", None) if db else None + if not callable(finder): + return None + try: + return finder( + source=source.platform.value, + user_id=source.user_id, + session_key=session_key, + chat_id=source.chat_id if allow_peer_fallback else None, + chat_type=source.chat_type if allow_peer_fallback else None, + thread_id=source.thread_id, + ) + except Exception as exc: + logger.debug("Gateway session DB recovery failed for %s: %s", session_key, exc) + if raise_on_lookup_error: + raise + return None + + def _recover_session_from_db( + self, + *, + session_key: str, + source: SessionSource, + now: datetime, + raise_on_lookup_error: bool = False, + ) -> Optional[SessionEntry]: + """Rebuild a missing session-key mapping from durable state.db data. + + Returns ``None`` when no row is recoverable, or when the recovered + session is already overdue under the reset policy — the row is then + durably promoted to a reset boundary instead of resurrected. + """ + entry, migrated_legacy = self._query_recoverable_row( + session_key=session_key, + source=source, + now=now, + raise_on_lookup_error=raise_on_lookup_error, + ) + if entry is None: + return None + reset_reason = self._should_reset(entry, source) + if reset_reason: + self._promote_session_reset( + session_key, entry.session_id, reset_reason, + log=lambda exc: logger.debug( + "Gateway recovered-session reset promotion failed for %s: %s", + session_key, exc, + ), + ) + return None + self._reopen_session_row(session_key, entry.session_id) + if migrated_legacy: + self._record_gateway_session_peer( + entry.session_id, session_key, source, display_name=entry.display_name, + ) + return entry + + def _query_recoverable_session(self, *, session_key, source, now): + """DB-only half of _recover_session_from_db (no lock needed). + + Returns a SessionEntry or None. Caller assigns _entries[key] under + lock. The row is NOT reopened here: the caller evaluates reset policy + first (an agent_close/ws_orphan row may need promotion to a real reset + boundary instead). + """ + entry, migrated_legacy = self._query_recoverable_row( + session_key=session_key, source=source, now=now, + ) + if entry is not None and migrated_legacy: + self._record_gateway_session_peer( + entry.session_id, session_key, source, display_name=entry.display_name, + ) + return entry + + def _query_recoverable_row( + self, *, session_key, source, now, raise_on_lookup_error=False, + ) -> tuple[Optional[SessionEntry], bool]: + """Find and gate a recoverable row -> (entry or None, migrated_legacy). + + The legacy (pre-workspace) Slack key fallback lives here: exact-key + lookup, claimed once per process; ``migrated_legacy`` tells the caller + to rewrite the peer row to the scoped key. + """ + legacy_key = self._legacy_slack_session_key(source) + recovered = self._find_gateway_session_row( + session_key=session_key, + source=source, + allow_peer_fallback=legacy_key is None, + 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) + ): + recovered = self._find_gateway_session_row( + session_key=legacy_key, + source=source, + allow_peer_fallback=False, + raise_on_lookup_error=raise_on_lookup_error, + ) + migrated_legacy = bool(recovered) + if not isinstance(recovered, dict): + return None, False + if not self._recovered_row_matches_source_scope(recovered, source): + return None, False + if not self._recovered_row_allowed_for_active_profile( + requested_session_key=session_key, + recovered=recovered, + ): + logger.warning( + "Gateway session DB recovery ignored %s for %s because " + "the row belongs to a different profile", + recovered.get("session_key"), + session_key, + ) + return None, False + entry = self._create_entry_from_recovered_row( + row=recovered, session_key=session_key, source=source, now=now, + ) + return entry, migrated_legacy + + def _promote_session_reset(self, session_key: str, session_id: str, reason: str, *, log) -> None: + """End *session_id* with *reason* via ``promote_to_session_reset``. + + Promote (not plain ``end_session``): a row already ended with a + recoverable accidental reason (agent_close / ws_orphan_reap) must be + upgraded to the explicit boundary, or stale-route recovery resurrects + it over the reset. Falls back to ``end_session`` on old SessionDBs. + ``log(exc)`` reports failures (each caller has its own message). + """ + try: + db = self._db_for_key(session_key) + promote = getattr(db, "promote_to_session_reset", None) + if callable(promote): + promote(session_id, reason) + else: + db.end_session(session_id, reason) + except Exception as exc: + log(exc) + + def _reopen_session_row(self, session_key: str, session_id: str, *, log_prefix: str = "") -> None: + """Best-effort ``reopen_session``; failures are debug-logged only.""" + try: + self._db_for_key(session_key).reopen_session(session_id) + except Exception as exc: + if log_prefix: + logger.debug("%s: %s", log_prefix, exc) + else: + logger.debug("Gateway session DB reopen failed for %s: %s", session_key, exc) + + def _record_gateway_session_peer( + self, + session_id: str, + session_key: str, + source: Optional[SessionSource], + display_name: Optional[str] = None, + include_compression_ancestors: bool = False, + ) -> None: + """Persist the routing peer for an existing gateway session row.""" + db = self._db_for_key(session_key) + if not db or not source: + return + recorder = getattr(db, "record_gateway_session_peer", None) + if not callable(recorder): + return + peer = dict( + source=source.platform.value, + user_id=source.user_id, + session_key=session_key, + chat_id=source.chat_id, + chat_type=source.chat_type, + 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, + include_compression_ancestors=include_compression_ancestors, + ) + except TypeError: + # Older SessionDB without display_name/origin_json kwargs. + try: + recorder(session_id, **peer) + except Exception as exc: + logger.debug("Gateway session peer record failed for %s: %s", session_key, exc) + except Exception as exc: + logger.debug("Gateway session peer record failed for %s: %s", session_key, exc) + + def _adopt_legacy_slack_entry(self, source: SessionSource, session_key: str) -> None: + """One-time migration of pre-workspace-scope Slack keys. + + MOVE (not copy) the legacy entry so a second workspace with identical + Slack ids cannot attach to the same transcript. Adopt when the legacy + origin names the same workspace; a scope-less DM is claimed once by + the first workspace; a scope-less channel/group is refused (channel + ids collide across workspaces). + """ + legacy_key = self._legacy_slack_session_key(source) + if not legacy_key: + return + migrated: Optional[SessionEntry] = None + with self._lock: + self._ensure_loaded_locked() + legacy_entry = self._entries.get(legacy_key) + if session_key not in self._entries and legacy_entry is not None: + origin_scope = getattr(legacy_entry.origin, "scope_id", None) + if origin_scope is not None: + adopt = origin_scope == source.scope_id + else: + adopt = source.chat_type == "dm" + if adopt and self._claim_legacy_slack_key(legacy_key): + migrated = self._entries.pop(legacy_key) + migrated.session_key = session_key + migrated.origin = source + migrated.platform = source.platform + migrated.chat_type = source.chat_type + self._entries[session_key] = migrated + if migrated is not None: + self._save_entries() + self._record_gateway_session_peer( + migrated.session_id, session_key, source, display_name=migrated.display_name, + ) + + def _finish_route_transition( + self, + session_key: str, + *, + end_session_id: Optional[str], + end_reason: str, + create_kwargs: Optional[Dict[str, Any]], + origin: Optional[SessionSource], + display_name: Optional[str], + during: str = "", + ) -> None: + """SQLite side of a routing transition, outside ``_lock``. + + Promotes the predecessor row to an explicit reset boundary (with the + specific reason so state.db is auditable, e.g. ``resume_pending_expired`` + vs a plain ``session_reset``), then INSERTs the new row + routing peer. + Both are best-effort: failures are warned and self-healed by the next + per-turn peer refresh. + """ + if self._db_for_key(session_key) and end_session_id: + self._promote_session_reset( + session_key, end_session_id, end_reason, + log=lambda e: logger.warning( + "Failed to end predecessor session row %s for %s%s: %s — " + "the old row remains open and may win restart recovery " + "until the next successful peer refresh", + end_session_id, session_key, during, e, + ), + ) + if self._db_for_key(session_key) and create_kwargs: + self._create_session_row( + session_key, create_kwargs, origin, display_name, + log=lambda e: logger.warning( + "Failed to create session row %s for %s%s: %s — deferring " + "to the self-healing peer refresh on the next turn", + create_kwargs.get("session_id"), session_key, during, e, + ), + ) + + @staticmethod + def _session_create_kwargs( + *, session_id, session_key, origin, source_value, display_name, parent_session_id, + ) -> Dict[str, Any]: + """kwargs for ``SessionDB.create_session``. + + 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, + "user_id": origin.user_id if origin else None, + "session_key": session_key, + "chat_id": origin.chat_id if origin else None, + "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, + "display_name": display_name, + "parent_session_id": parent_session_id, + "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: + """INSERT a session row and record its routing peer; ``log(exc)`` on failure. + + A failed create is a routing hazard (visible warning), but the row is + self-healed with full identity by the next per-turn peer refresh. + """ + 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, + ) + except Exception as e: + log(e) diff --git a/gateway/session_transcript.py b/gateway/session_transcript.py new file mode 100644 index 0000000000..e2f4430451 --- /dev/null +++ b/gateway/session_transcript.py @@ -0,0 +1,641 @@ +"""SessionStore transcript I/O: SQLite append with a per-session retry queue, +compression-reroute following, FTS corruption recovery, rewrite/rewind/load. + +Mixin split out of ``gateway/session.py``; bound onto ``SessionStore`` via the MRO. +""" + +from __future__ import annotations + +import logging +import threading +from agent.turn_context import extract_api_content_sidecar +from typing import TYPE_CHECKING, Any, Dict, List, Optional + +if TYPE_CHECKING: + from gateway.session import SessionEntry + +# Log-record parity with the origin module. +logger = logging.getLogger("gateway.session") + + +def _plain_text(content) -> str: + """Text of a message content (str or text-part list); "" for anything else.""" + if isinstance(content, list): + parts = [p.get("text", "") for p in content if isinstance(p, dict) and p.get("type") == "text"] + return "\n".join(t for t in parts if t) + return content if isinstance(content, str) else "" + + +class SessionTranscriptMixin: + """SessionStore transcript I/O: SQLite append with a per-session retry queue, + compression-reroute following, FTS corruption recovery, + rewrite/rewind/load. + """ + + def _compression_tip_for_session_id(self, session_id: Optional[str]) -> Optional[str]: + """Latest compression continuation for *session_id* (heals a mapping + left pointing at a compressed parent by a restart or failed send).""" + if not session_id: + return session_id + db = self._db_for_session_id(session_id) + if db is None: + return session_id + try: + return db.get_compression_tip(session_id) or session_id + except Exception: + logger.debug("Compression-tip lookup failed for session %s", session_id, exc_info=True) + return session_id + + def _heal_compression_tip_locked( + self, + entry: "SessionEntry", + original_session_id: Optional[str], + canonical_session_id: Optional[str], + ) -> bool: + """Rewrite *entry* to the compression continuation if stale. Lock held.""" + if ( + not original_session_id + or not canonical_session_id + or entry.session_id != original_session_id + or canonical_session_id == original_session_id + ): + return False + logger.info( + "SessionStore healed compressed session mapping: %s -> %s", + entry.session_id, + canonical_session_id, + ) + entry.session_id = canonical_session_id + return True + + def advance_compression_session( + self, + session_key: str, + expected_session_id: str, + target_session_id: str, + ) -> Optional[SessionEntry]: + """CAS-advance one route along an already-verified compression lineage. + + Unlike ``switch_session`` this never ends/reopens SQLite rows (the + compression transaction owns that). ``None`` means the route moved + after the caller's snapshot (e.g. /new) — caller must fail closed. + """ + if not session_key or not expected_session_id or not target_session_id: + return None + + with self._lock: + entry = self._entry_locked(session_key) + if entry is None: + return None + if entry.session_id == target_session_id: + 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, + ): + return None + # Bookkeeping, not user activity: leave ``updated_at`` alone. + self._save() + return entry + + def _get_transcript_drain_lock(self): + """Return the lock that serializes pending-queue drain boundaries.""" + return self._lazy("_transcript_drain_lock", threading.RLock) + + def append_to_transcript(self, session_id: str, message: Dict[str, Any], skip_db: bool = False) -> None: + """Serialize transcript draining across queue migration boundaries.""" + 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 + ) + + def _follow_reroutes(self, session_id: str) -> str: + """Follow the compression reroute chain (cycle-guarded).""" + reroutes = self._lazy("_transcript_reroutes", dict) + seen = set() + while session_id in reroutes and session_id not in seen: + seen.add(session_id) + 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. + """ + 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)) + if spool_path is not None: + self._lazy("_spooled_drop_sessions", set).add(session_id) + logger.warning( + "Session DB transcript pending queue full for %s " + "(cap=%d); spooled oldest message to %s for replay " + "after DB recovery", + session_id, self._MAX_PENDING_PER_SESSION, spool_path, + ) + else: + logger.warning( + "Session DB transcript pending queue full for %s " + "(cap=%d); dropping oldest message to make room " + "(on-disk spool unavailable)", + session_id, self._MAX_PENDING_PER_SESSION, + ) + return pending + + def _divert_transcript_after_db_replaced( + self, session_id: str, queue_session_id: str, exc: Exception + ) -> None: + """Stop SQLite writes on a replaced/quarantined handle and divert the backlog. + + Retrying cannot succeed and the FTS rebuild must never run on this + handle; the pending queue goes to the on-disk spool + JSONL fallback. + """ + logger.error( + "Session DB refused further writes on this handle for " + "%s (%s); stopping SQLite writes and diverting pending " + "transcripts to the on-disk fallback: %s", + session_id, type(exc).__name__, exc, + ) + with self._transcript_retry_lock: + remaining = list(self._dirty_transcripts.get(queue_session_id, [])) + self._dirty_transcripts.pop(queue_session_id, None) + self._transcript_append_failures.pop(session_id, None) + for dropped in remaining: + try: + from gateway.shutdown_flush import spool_dropped_transcript_message + + spool_dropped_transcript_message(session_id, dropped) + except Exception: + logger.warning( + "pending fallback failed for replaced " + "state.db transcript on %s", + session_id, + exc_info=True, + ) + try: + from hermes_state import divert_session_transcript_jsonl + divert_session_transcript_jsonl(session_id, remaining) + except Exception: + logger.warning( + "JSONL divert failed for replaced state.db " + "transcript on %s", + session_id, + exc_info=True, + ) + + def _live_compression_child(self, session_id: str) -> str: + """Transitive compression tip of *session_id* if it is a different, still-live + row, else "" (a depth-1 lookup misses multi-hop lineages). + + Uses the PARENT's proven owner handle: the child's id is not published + until after its write succeeds, so a by-id lookup would fall back to + the ambient store. + """ + owner_db = self._db_for_session_id(session_id) + if owner_db is None: + return "" + tip = owner_db.get_compression_tip(session_id) + if tip and tip != session_id: + tip_row = owner_db.get_session(tip) + if tip_row is not None and tip_row.get("ended_at") is None: + return str(tip) + return "" + + def _migrate_transcript_queue_to_child( + self, session_id: str, queue_session_id: str, child_id: str, pending: list, msg + ) -> list: + """Move the retry queue + failure counter from parent to child and publish + the reroute (retry lock held). Returns the child's pending list. + + Older parent backlog must precede messages already queued directly on + the child. Routing is published only AFTER the queue moved (caller), so + new child writes cannot bypass older parent backlog. + """ + if pending and pending[0] is msg: + pending.pop(0) + existing_child_pending = self._dirty_transcripts.get(child_id, []) + if pending: + pending.extend(existing_child_pending) + self._dirty_transcripts[child_id] = pending + elif existing_child_pending: + pending = existing_child_pending + self._dirty_transcripts.pop(queue_session_id, None) + 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), + ) + self._transcript_reroutes[session_id] = child_id + return pending + + def _publish_transcript_reroute(self, session_id: str, child_id: str) -> None: + """Repoint every route at the compression child and save (index authoritative again).""" + with self._lock: + for entry in self._entries.values(): + if entry.session_id == session_id: + entry.session_id = child_id + self._save() + _hints = getattr(self, "_session_owner_hints", None) + if _hints: + _hints.pop(child_id, None) + + def _append_to_transcript_serialized( + self, session_id: str, message: Dict[str, Any] + ) -> None: + """Append a message to a session's transcript (SQLite), draining the + per-session retry queue.""" + with self._transcript_retry_lock: + pending = self._enqueue_transcript_message(session_id, message) + msg = pending[0] + queue_session_id = session_id + + def _ack_head() -> bool: + """Pop the acknowledged head (retry lock held). True if queue drained.""" + if pending and pending[0] is msg: + pending.pop(0) + if not pending: + self._dirty_transcripts.pop(queue_session_id, None) + self._transcript_append_failures.pop(session_id, None) + return True + return False + + # DB write outside the retry lock so other sessions can append. + while True: + try: + self._append_transcript_message(session_id, msg) + except Exception as exc: + from hermes_state import ( + CompressionSessionClosedError, + StateDbCorruptError, + StateDbReplacedError, + ) + + if isinstance(exc, (StateDbReplacedError, StateDbCorruptError)): + self._divert_transcript_after_db_replaced(session_id, queue_session_id, exc) + return + + if isinstance(exc, CompressionSessionClosedError): + # Adopt only a different, still-live compression tip, else + # fail closed. + _owner_key = self._owner_key_for_session_id(session_id) + child_id = self._live_compression_child(session_id) + if child_id: + # Record the child's owner BEFORE writing to it (the + # reroute is published only after the write succeeds + # — load-bearing for backlog order). + if _owner_key: + self._lazy("_session_owner_hints", dict)[child_id] = _owner_key + try: + self._append_transcript_message(child_id, msg) + except Exception as reroute_exc: + exc = reroute_exc + else: + with self._transcript_retry_lock: + pending = self._migrate_transcript_queue_to_child( + session_id, queue_session_id, child_id, pending, msg + ) + queue_session_id = child_id + self._publish_transcript_reroute(session_id, child_id) + if not pending: + return + msg = pending[0] + session_id = child_id + continue + else: + # Permanent routing invariant failure, not a transient + # outage: drop it so it cannot poison later writes. + with self._transcript_retry_lock: + _ack_head() + logger.error( + "Session DB transcript append rejected for compression-ended " + "%s with no unique live child; not retrying", + session_id, + ) + return + if self._is_fts_corruption_error(exc) and self._rebuild_fts_once(): + try: + self._append_transcript_message(session_id, msg) + except Exception as retry_exc: + exc = retry_exc + else: + with self._transcript_retry_lock: + _ack_head() + continue + with self._transcript_retry_lock: + failures = self._transcript_append_failures.get(session_id, 0) + 1 + self._transcript_append_failures[session_id] = failures + logger.warning( + "Session DB transcript append failed for %s " + "(failure_count=%d, pending=%d); will retry: %s", + session_id, failures, len(pending), exc, + ) + return + else: + with self._transcript_retry_lock: + queue_empty = _ack_head() + if not queue_empty: + msg = pending[0] + if queue_empty: + # Backlog clear: replay cap-dropped messages spooled to disk. + self._drain_spooled_drops(session_id) + return + continue + + def _drain_spooled_drops(self, session_id: str) -> None: + """Replay cap-dropped spooled transcript messages after DB recovery. + + Best-effort: replay failures keep the spool files for the next + successful flush; nothing here may raise into the caller. + """ + spooled_sessions = getattr(self, "_spooled_drop_sessions", None) + if not spooled_sessions or session_id not in spooled_sessions: + return + try: + from gateway.shutdown_flush import drain_transcript_spool + + _replayed, remaining = drain_transcript_spool( + session_id, + lambda message: self._append_transcript_message( + session_id, message + ), + ) + if not remaining: + spooled_sessions.discard(session_id) + except Exception as exc: + logger.warning("Failed to drain transcript spool for %s: %s", session_id, exc) + + def _append_transcript_message(self, session_id: str, message: Dict[str, Any]) -> None: + """Write one transcript row. Caller handles retry queuing.""" + _db = self._db_for_session_id(session_id) + if _db is None: + # Named profile with no resolvable home yet: defer (caller queues) + # instead of writing into the ambient store. + raise RuntimeError( + f"no owning session store for {session_id}; deferring transcript write" + ) + is_assistant = message.get("role") == "assistant" + _db.append_message( + session_id=session_id, + role=message.get("role", "unknown"), + content=message.get("content"), + tool_name=message.get("tool_name"), + tool_calls=message.get("tool_calls"), + tool_call_id=message.get("tool_call_id"), + reasoning=message.get("reasoning") if is_assistant else None, + reasoning_content=message.get("reasoning_content") if is_assistant else None, + reasoning_details=message.get("reasoning_details") if is_assistant else None, + codex_reasoning_items=message.get("codex_reasoning_items") if is_assistant else None, + codex_message_items=message.get("codex_message_items") if is_assistant else None, + platform_message_id=(message.get("platform_message_id") or message.get("message_id")), + observed=bool(message.get("observed")), + timestamp=message.get("timestamp"), + # Exact bytes sent to the API (prompt-cache-stable replay); must + # survive every persistence path or the next replay diverges. + api_content=extract_api_content_sidecar(message), + # Presentation typing (e.g. "internal_notification"); DB-only. + display_kind=message.get("display_kind"), + 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. + + A bare SQLITE_CORRUPT can mean structural B-tree damage; only errors + naming ``messages_fts`` or carrying FTS provenance (per + ``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: + return True + import sqlite3 + + from hermes_state import SessionDB + + if isinstance(exc, sqlite3.DatabaseError): + return SessionDB._is_fts_write_corruption_error(exc) + return False + + def _rebuild_fts_once(self) -> bool: + """Attempt FTS5 ``rebuild`` once per store lifetime; True if any index was rebuilt.""" + if self._fts_rebuild_attempted: + return False + self._fts_rebuild_attempted = True + db = self._db + if db is None or not hasattr(db, "rebuild_fts"): + return False + # WAL split-brain guard: skip when a foreign process holds state.db. + if hasattr(db, "_foreign_state_db_holders"): + foreign_holders = db._foreign_state_db_holders() + if foreign_holders: + logger.warning( + "Skipping Session DB FTS rebuild while foreign processes " + "hold the database or WAL sidecars (%s); canonical " + "transcript writes remain available.", + foreign_holders, + ) + return False + try: + rebuilt = db.rebuild_fts() + except Exception as exc: + 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, + ) + return rebuilt > 0 + + def _clear_dirty_transcript(self, session_id: str) -> None: + """Drop queued pending messages so a rewrite/rewind doesn't re-insert them.""" + with self._transcript_retry_lock: + self._dirty_transcripts.pop(session_id, None) + self._transcript_append_failures.pop(session_id, None) + + def has_platform_message_id( + self, session_id: str, platform_message_id: str + ) -> bool: + """Whether a message with this platform_message_id is persisted (False without a DB).""" + db = self._db_for_session_id(session_id) + if not db: + return False + try: + 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 + + def rewrite_transcript( + self, + session_id: str, + messages: List[Dict[str, Any]], + active_only: bool = False, + reject_active_turn_lease: bool = False, + ) -> bool: + """Replace a session's transcript (/retry, /compress). + + DESTRUCTIVE by default: ``active_only=False`` DELETEs every row incl. + soft-archived compaction history (pass ``active_only=True`` for sessions + that may carry archived rows). True when the write lands or there is no + DB, False on failure — callers committing a destructive change on top + (/compress repointing) must check it. ``reject_active_turn_lease`` is + for user-initiated rewrites that do not own the cross-process turn lease. + """ + db = self._db_for_session_id(session_id) + if not db: + return True + with self._get_transcript_drain_lock(): + try: + db.replace_messages( + session_id, + messages, + active_only=active_only, + reject_active_turn_lease=reject_active_turn_lease, + ) + except Exception as e: + logger.debug("Failed to rewrite transcript in DB: %s", e) + return False + self._clear_dirty_transcript(session_id) + return True + + def load_transcript(self, session_id: str) -> List[Dict[str, Any]]: + """Load all messages from a session's transcript (state.db is canonical). + + Reads follow the same routing writes use: the in-memory reroute map + (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) + try: + # Durable successor survives restart; the reroute map doesn't. + tip = self._db_for_session_id(session_id).get_compression_tip(session_id) + if tip: + session_id = tip + except Exception: + pass + try: + # repair_alternation: this feeds LIVE REPLAY; heal a durable + # user;user wedge once here instead of on every request. + return self._db_for_session_id(session_id).get_messages_as_conversation( + session_id, repair_alternation=True + ) + except Exception as e: + # Empty history is valid data; a failed canonical read is not — + # live-replay callers must fail closed, not start from []. + logger.error( + "Transcript read failed for session %s; refusing to treat the " + "conversation as empty: %s", + session_id, + e, + exc_info=True, + ) + raise TranscriptReadError(session_id) from e + + def rewind_session( + self, + session_id: str, + n: int = 1, + *, + require_retryable_composite: bool = False, + ) -> Optional[Dict[str, Any]]: + """Back up ``n`` user turns via soft-delete (``active=0``), mirroring CLI ``/undo [N]``. + + Returns ``{"rewound_count", "turns_undone", "target_text"}`` or ``None`` + (no DB / no user turn). ``n`` clamps to the oldest user turn. + ``require_retryable_composite`` is the gateway ``/retry`` guard: the + selected turn must be a composite carrier whose live payload is + losslessly replayable as text before anything changes. + """ + db = self._db_for_session_id(session_id) + if not db: + return None + with self._get_transcript_drain_lock(): + if n < 1: + n = 1 + from agent.context_compressor import ( + retryable_user_text, + split_user_originated_turn, + user_originated_turn_view, + ) + + try: + expected_active_ids = db.get_active_message_ids(session_id) + durable = db.get_messages_as_conversation( + session_id, + include_row_ids=True, + ) + user_indices = [ + index + for index, message in enumerate(durable) + if user_originated_turn_view(message) is not None + ] + if not user_indices: + return None + turns_undone = min(n, len(user_indices)) + target = durable[user_indices[-turns_undone]] + target_id = target.get("_row_id") + if not isinstance(target_id, int): + return None + handoff, target_view = split_user_originated_turn(target) + if target_view is None: + return None + if require_retryable_composite and handoff is None: + return None + except Exception as e: + logger.debug("rewind_session: failed to resolve canonical target: %s", e) + return None + if require_retryable_composite: + # Keep replay-policy failures distinct from persistence errors + # so /retry can explain why the selected carrier is unsafe. + target_text = retryable_user_text(target_view.get("content")) + try: + result = db.rewind_to_message( + session_id, + target_id, + preserve_compaction_handoff=handoff is not None, + expected_active_ids=expected_active_ids, + expected_target_content=target_view.get("content"), + ) + except ValueError as e: + logger.debug("rewind_session: %s", e) + return None + except Exception as e: + logger.debug("rewind_session: rewind_to_message failed: %s", e) + return None + self._clear_dirty_transcript(session_id) + # ``target_view`` is the live projection; a composite carrier's raw + # row holds the summary wrapper and must not be echoed as prompt. + if not require_retryable_composite: + target_text = _plain_text(target_view.get("content") or "") + return { + "rewound_count": result.get("rewound_count", 0), + "turns_undone": turns_undone, + "target_text": target_text, + }