refactor(gateway): session mixins — compact docstrings/comments, fold short call sites (AST-identical)

This commit is contained in:
Teknium
2026-09-02 20:52:08 -07:00
parent c4655a638c
commit f6c2c1a80a
4 changed files with 325 additions and 599 deletions
+87 -152
View File
@@ -1,8 +1,6 @@
"""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.
"""
"""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
@@ -35,11 +33,10 @@ def _plain_text(content) -> str:
def _spool_dropped(session_id: str, message: Dict[str, Any]):
"""Spool one evicted/undeliverable message to disk (same machinery as the
shutdown flush, so it is replayed after DB recovery); path or None."""
"""Spool one evicted/undeliverable message to disk (same machinery as the shutdown flush, so it
is replayed after DB recovery); path or None."""
try:
from gateway.shutdown_flush import spool_dropped_transcript_message
return spool_dropped_transcript_message(session_id, message)
except Exception:
return None
@@ -52,8 +49,8 @@ class SessionTranscriptMixin:
_MAX_PENDING_PER_SESSION = 200 # in-memory pending messages per session (DB broken)
def _compression_tip_for_session_id(self, session_id: Optional[str]) -> Optional[str]:
"""Latest compression continuation for *session_id* (heals a mapping
left pointing at a compressed parent by a restart or failed send)."""
"""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)
@@ -67,8 +64,7 @@ class SessionTranscriptMixin:
def _heal_compression_tip_locked(
self, entry: "SessionEntry", original_session_id: Optional[str],
canonical_session_id: Optional[str],
) -> bool:
canonical_session_id: Optional[str]) -> bool:
"""Rewrite *entry* to the compression continuation if stale. Lock held."""
if (
not original_session_id
@@ -78,24 +74,20 @@ class SessionTranscriptMixin:
):
return False
logger.info(
"SessionStore healed compressed session mapping: %s -> %s",
entry.session_id, canonical_session_id,
)
"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.
"""
"""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:
@@ -106,8 +98,7 @@ class SessionTranscriptMixin:
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()
self._save() # bookkeeping, not user activity: leave ``updated_at`` alone
return entry
def _get_transcript_drain_lock(self):
@@ -139,34 +130,25 @@ class SessionTranscriptMixin:
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,
)
"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,
)
"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.
"""
"""Stop SQLite writes on a replaced/quarantined handle and divert the backlog to the on-disk
spool + JSONL fallback: retrying cannot succeed and the FTS rebuild must never run here."""
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,
)
"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)
@@ -174,30 +156,24 @@ class SessionTranscriptMixin:
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,
)
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,
)
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.
"""
"""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 unpublished until its write succeeds, so a by-id lookup would hit the ambient
store."""
owner_db = self._db_for_session_id(session_id)
if owner_db is None:
return ""
@@ -211,13 +187,10 @@ class SessionTranscriptMixin:
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.
"""
"""Move the retry queue + failure counter from parent to child and record 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, [])
@@ -230,8 +203,7 @@ class SessionTranscriptMixin:
previous_failures = self._transcript_append_failures.pop(queue_session_id, 0)
if previous_failures:
self._transcript_append_failures[child_id] = max(
previous_failures, self._transcript_append_failures.get(child_id, 0),
)
previous_failures, self._transcript_append_failures.get(child_id, 0))
self._transcript_reroutes[session_id] = child_id
return pending
@@ -247,8 +219,8 @@ class SessionTranscriptMixin:
_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."""
"""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]
@@ -272,11 +244,9 @@ class SessionTranscriptMixin:
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)
@@ -293,8 +263,7 @@ class SessionTranscriptMixin:
else:
with self._transcript_retry_lock:
pending = self._migrate_transcript_queue_to_child(
session_id, queue_session_id, child_id, pending, msg
)
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:
@@ -303,15 +272,13 @@ class SessionTranscriptMixin:
session_id = child_id
continue
else:
# Permanent routing invariant failure, not a transient
# outage: drop it so it cannot poison later writes.
# 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,
)
"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:
@@ -326,10 +293,8 @@ class SessionTranscriptMixin:
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,
)
"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:
@@ -343,17 +308,13 @@ class SessionTranscriptMixin:
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.
"""
"""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."""
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),
)
@@ -366,11 +327,10 @@ class SessionTranscriptMixin:
"""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.
# 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"
)
f"no owning session store for {session_id}; deferring transcript write")
is_assistant = message.get("role") == "assistant"
_db.append_message(
session_id=session_id,
@@ -387,8 +347,8 @@ class SessionTranscriptMixin:
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.
# 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"),
@@ -397,19 +357,14 @@ class SessionTranscriptMixin:
@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.
"""
"""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 (``SessionDB._is_fts_write_corruption_error``) may authorize the one-shot
rebuild-and-retry; everything else takes the retry path."""
if "messages_fts" in str(exc).lower():
return True
import sqlite3
from hermes_state import SessionDB
return isinstance(exc, sqlite3.DatabaseError) and SessionDB._is_fts_write_corruption_error(exc)
def _rebuild_fts_once(self) -> bool:
@@ -425,11 +380,9 @@ class SessionTranscriptMixin:
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,
)
"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()
@@ -459,17 +412,13 @@ class SessionTranscriptMixin:
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.
"""
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
@@ -477,8 +426,7 @@ class SessionTranscriptMixin:
try:
db.replace_messages(
session_id, messages, active_only=active_only,
reject_active_turn_lease=reject_active_turn_lease,
)
reject_active_turn_lease=reject_active_turn_lease)
except Exception as e:
logger.debug("Failed to rewrite transcript in DB: %s", e)
return False
@@ -486,12 +434,9 @@ class SessionTranscriptMixin:
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.
"""
"""Load all messages from a session's transcript (state.db is canonical). Reads follow the
same routing writes use — the in-memory reroute map, then the durable compression tip —
otherwise the transcript "vanishes" while every message sits under the child."""
if not self._db_for_session_id(session_id):
return []
session_id = self._follow_reroutes(session_id)
@@ -503,32 +448,25 @@ class SessionTranscriptMixin:
except Exception:
pass
try:
# repair_alternation: this feeds LIVE REPLAY; heal a durable
# user;user wedge once here instead of on every request.
# repair_alternation: this feeds LIVE REPLAY; heal a durable user;user wedge once here.
return self._db_for_session_id(session_id).get_messages_as_conversation(
session_id, repair_alternation=True
)
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 [].
# 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,
)
"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.
"""
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
@@ -537,13 +475,11 @@ class SessionTranscriptMixin:
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)
index for index, message in enumerate(durable)
if user_originated_turn_view(message) is not None
]
if not user_indices:
@@ -562,15 +498,14 @@ class SessionTranscriptMixin:
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.
# 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"),
)
expected_target_content=target_view.get("content"))
except ValueError as e:
logger.debug("rewind_session: %s", e)
return None
@@ -578,8 +513,8 @@ class SessionTranscriptMixin:
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.
# ``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 {