refactor(gateway/session): split SessionStore into persistence/recovery/lifecycle/transcript mixins by call-graph cohesion; compact wire helpers

This commit is contained in:
Teknium
2026-09-02 16:15:42 -07:00
parent c89819bec2
commit d7bdf2788d
5 changed files with 2264 additions and 2144 deletions

File diff suppressed because it is too large Load Diff

View File

@@ -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

View File

@@ -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.<name>`` 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:<platform>:...) 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()

519
gateway/session_recovery.py Normal file
View File

@@ -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)

View File

@@ -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,
}