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