refactor(gateway/session): split SessionStore into persistence/recovery/lifecycle/transcript mixins by call-graph cohesion; compact wire helpers
This commit is contained in:
2207
gateway/session.py
2207
gateway/session.py
File diff suppressed because it is too large
Load Diff
385
gateway/session_lifecycle.py
Normal file
385
gateway/session_lifecycle.py
Normal 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
|
||||
656
gateway/session_persistence.py
Normal file
656
gateway/session_persistence.py
Normal 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
519
gateway/session_recovery.py
Normal 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)
|
||||
641
gateway/session_transcript.py
Normal file
641
gateway/session_transcript.py
Normal 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,
|
||||
}
|
||||
Reference in New Issue
Block a user