refactor(agent/E_session): by-hand docstring/comment compaction across 9 files; transcript_repair row lookup helper; activity snapshot literal

This commit is contained in:
Teknium
2026-09-02 22:23:45 -07:00
parent 04a2c422f6
commit 63f1950b4f
9 changed files with 209 additions and 305 deletions
+54 -69
View File
@@ -1,8 +1,6 @@
"""Durable transcript persistence for ``AIAgent`` (mixin; MRO-resolved from ``run_agent``).
SQLite session flush with intrinsic ``_DB_PERSISTED_MARKER`` dedup, ephemeral-scaffolding
filtering, the optional JSON session log and trajectory export.
"""
"""Durable transcript persistence for ``AIAgent`` (mixin; MRO-resolved from ``run_agent``): SQLite flush
with intrinsic ``_DB_PERSISTED_MARKER`` dedup, ephemeral-scaffolding filtering, optional JSON session log,
trajectory export."""
import hashlib
import json
import logging
@@ -25,9 +23,7 @@ from agent.trajectory import convert_scratchpad_to_think, save_trajectory as _sa
from agent.transcript_repair import sync_flushed_message_markers
from utils import atomic_json_write
# Same logger name as the origin module so log records / caplog filters are unchanged.
logger = logging.getLogger("run_agent")
logger = logging.getLogger("run_agent") # origin module's name: log records / caplog filters unchanged
# Flags marking ephemeral recovery scaffolding the loop pops before appending the real response.
# Persistence must skip them or a resumed session replays synthetic turns / breaks prefix-cache reuse.
@@ -51,16 +47,15 @@ def _is_ephemeral_scaffolding(msg: Any) -> bool:
return isinstance(msg, dict) and any(msg.get(flag) for flag in _EPHEMERAL_SCAFFOLDING_FLAGS)
# `_DB_PERSISTED_MARKER` (agent.context_compressor) is the intrinsic "already written to SQLite" marker:
# an id(msg) set can alias a freed dict's address onto a new message, a key on the dict cannot. The `_`
# prefix is mandatory (wire sanitizers strip `_` keys). CONTRACT: the marker asserts the dict's CONTENT
# is durable as written — any in-place mutation that must persist MUST pop it (turn_finalizer,
# context_compressor).
# `_DB_PERSISTED_MARKER` (agent.context_compressor) is the intrinsic "already written to SQLite" marker: an
# id(msg) set can alias a freed dict's address onto a new message, a key on the dict cannot. The `_` prefix is
# mandatory (wire sanitizers strip `_` keys). CONTRACT: the marker asserts the dict's CONTENT is durable as
# written — any in-place mutation that must persist MUST pop it (turn_finalizer, context_compressor).
def _safe_session_filename_component(session_id: str) -> str:
"""Path-safe filename component for a (possibly untrusted ``X-Hermes-Session-Id``) session ID:
non ``[A-Za-z0-9_-]`` → ``_``, capped, plus a content hash when changed so distinct IDs cannot collide."""
"""Path-safe component for a (possibly untrusted ``X-Hermes-Session-Id``) ID: non ``[A-Za-z0-9_-]`` → ``_``,
capped, plus a content hash when changed so distinct IDs cannot collide."""
raw = str(session_id or "").strip()
sanitized = re.sub(r"[^\w-]", "_", raw).strip("._")[:96] or "session"
if raw and sanitized == raw:
@@ -69,9 +64,9 @@ def _safe_session_filename_component(session_id: str) -> str:
def _override_replaces_content(msg: Dict, content: Any, override: Any) -> bool:
"""Whether the persist user-message override may replace ``content``: a plain-text override must
not replace native image/audio blocks (a list override is the clean multimodal payload and does),
and never a message MERGED with a compaction summary (overwriting would drop the summary)."""
"""May the persist override replace ``content``? A plain-text override must not replace native image/audio
blocks (a list override is the clean multimodal payload and does), nor a message MERGED with a compaction
summary (overwriting would drop the summary)."""
return (
override is not None
and not msg.get(COMPRESSED_SUMMARY_METADATA_KEY)
@@ -80,8 +75,8 @@ def _override_replaces_content(msg: Dict, content: Any, override: Any) -> bool:
def _summary_display_kind(msg: Dict) -> Any:
"""Standalone handoffs are hidden so they never occupy the active user slot in retry/undo
dispatch; merge-into-tail carriers keep their prior visibility."""
"""Standalone handoffs are hidden so they never occupy the active user slot in retry/undo dispatch;
merge-into-tail carriers keep their prior visibility."""
if (
msg.get(COMPRESSED_SUMMARY_METADATA_KEY)
and user_originated_turn_view(msg) is None
@@ -123,10 +118,9 @@ def _persist_lock(agent):
# --- flush phases (module-level so the flush also works bound onto duck-typed agents) ---
def _db_flush_seed_ids(agent) -> set:
"""One-shot ``_flushed_db_message_ids`` seed (same session, after a non-empty flush); the scan
translates it to markers and the flush clears it."""
"""One-shot ``_flushed_db_message_ids`` seed (same session, after a non-empty flush); the scan translates
it to markers and the flush clears it."""
current_session_id = getattr(agent, "session_id", None)
seed_ids = None
if getattr(agent, "_flushed_db_message_session_id", None) == current_session_id and agent._last_flushed_db_idx != 0:
@@ -149,15 +143,13 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any
"""Build the session-db row for ``msg``, applying the persist override to THIS row only."""
role = msg.get("role", "unknown")
content = msg.get("content")
# api_content sidecar: exact bytes sent to the API when they differ from clean content, so
# replay reproduces the sent prefix byte-for-byte.
# api_content sidecar: exact bytes sent to the API when they differ from clean content (replay parity).
api_content = msg.get("api_content") if isinstance(msg.get("api_content"), str) else None
timestamp = msg.get("timestamp")
if is_current_turn_user and msg.get("role") == "user":
override = getattr(agent, "_persist_user_message_override", None)
if _override_replaces_content(msg, content, override):
# Live content is what the wire sent, the override is the clean transcript; keep the
# sent bytes in api_content so replay matches the wire.
# Live content is what the wire sent, the override is the clean transcript; keep the sent bytes.
if api_content is None and isinstance(content, str) and content != override:
api_content = content
content = override
@@ -166,8 +158,8 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any
timestamp = ov_timestamp
if api_content == content:
api_content = None
# get_messages_as_conversation replays rows through sanitize_context().strip(); capture the
# sent bytes when they would differ (compared in wire form).
# get_messages_as_conversation replays rows through sanitize_context().strip(); capture the sent bytes
# when they would differ (compared in wire form).
if (
api_content is None and role in ("user", "assistant") and isinstance(content, str) and content
and sanitize_context(content).strip() != content.strip()
@@ -200,15 +192,15 @@ def _db_flush_collect(agent, messages: List[Dict], conversation_history: Optiona
seed_ids = _db_flush_seed_ids(agent)
history_ids = {id(item) for item in (conversation_history or []) if isinstance(item, dict)}
ov_idx = getattr(agent, "_persist_user_message_idx", None)
# Also match the staged CLI dict by identity — the close safety-net may flush a shortened
# snapshot whose turn index refers to the full history.
# Also match the staged CLI dict by identity — the close safety-net may flush a shortened snapshot whose
# turn index refers to the full history.
pending_cli_message = getattr(agent, "_pending_cli_user_message", None)
batch_rows: List[Dict[str, Any]] = []
batch_msgs: List[Dict] = []
for msg_idx in range(_db_flush_scan_start(agent, messages), len(messages)):
msg = messages[msg_idx]
# The flush is append-only: a mid-turn persist of scaffolding could commit a synthetic
# turn the end-of-turn drop cannot un-write. Skip regardless of position.
# Append-only flush: a mid-turn persist of scaffolding would commit a synthetic turn the end-of-turn
# drop cannot un-write. Skip regardless of position.
if not isinstance(msg, dict) or _is_ephemeral_scaffolding(msg) or msg.get(_DB_PERSISTED_MARKER):
continue
# Already durable (history copy or caller-seeded): stamp so future flushes skip it.
@@ -235,8 +227,8 @@ def _db_flush_write(agent, batch_rows: List[Dict[str, Any]], batch_msgs: List[Di
def _db_flush_adopt_compression_tip(agent) -> bool:
"""Adopt the live continuation of a session closed by compression. Same-id tip = no continuation;
a tip whose row is missing or already ended is not adopted either."""
"""Adopt the live continuation of a compression-closed session. Same-id tip = no continuation; a tip
whose row is missing or already ended is not adopted either."""
old_id = agent.session_id
try:
tip = agent._session_db.get_compression_tip(old_id)
@@ -261,10 +253,9 @@ def _db_flush_adopt_compression_tip(agent) -> bool:
def _db_flush_failed(agent, e: Exception, batch_rows: List[Dict[str, Any]], adoption_budget: int) -> bool:
"""Classify a failed flush; True when the caller should retry once on an adopted compression tip."""
# Force a full re-scan next flush: an exception mid-loop leaves mixed dispositions.
agent._db_flush_scan_prefix = None
# The only place the SQLite error is visible before it becomes a bare False — classify it so
# the turn-end explanation can distinguish lock contention from disk-full/read-only.
agent._db_flush_scan_prefix = None # full re-scan next flush: an exception mid-loop leaves mixed dispositions
# The only place the SQLite error is visible before it becomes a bare False — classify it so the turn-end
# explanation can distinguish lock contention from disk-full/read-only.
from hermes_state import (
CompressionSessionClosedError,
StateDbCorruptError,
@@ -284,19 +275,17 @@ def _db_flush_failed(agent, e: Exception, batch_rows: List[Dict[str, Any]], adop
agent._last_persistence_error_cause, getattr(agent, "session_id", None), exc_info=True,
)
if isinstance(e, CompressionSessionClosedError):
# Compression race: another path rotated this session mid-write. Retry exactly once on the
# live tip; a second closed-parent write fails closed.
# Compression race: another path rotated this session mid-write. Retry exactly once on the live tip; a
# second closed-parent write fails closed.
if adoption_budget > 0 and _db_flush_adopt_compression_tip(agent):
return True
# The flag lets the turn explanation name compression rotation instead of misleading
# full-disk advice.
agent._compression_adoption_failed = True
agent._compression_adoption_failed = True # lets the turn explanation name rotation, not full-disk advice
logger.warning("Session DB append_message failed: %s", e)
return False
def _session_log_entry(agent, msg: Dict[str, Any]) -> Dict[str, Any]:
"""Copy of ``msg`` with scratchpad tags normalised and credentials redacted (respects HERMES_REDACT_SECRETS)."""
"""Copy of ``msg`` with scratchpad tags normalised and credentials redacted (honours HERMES_REDACT_SECRETS)."""
if "content" not in msg:
return msg
content = msg["content"]
@@ -306,8 +295,8 @@ def _session_log_entry(agent, msg: Dict[str, Any]) -> Dict[str, Any]:
def _existing_log_is_larger(log_file, count: int) -> bool:
"""Never overwrite a larger log with fewer messages (resumed agent with partial history);
a corrupted existing file allows the overwrite."""
"""Never overwrite a larger log with fewer messages (resumed agent with partial history); a corrupted
existing file allows the overwrite."""
if not log_file.exists():
return False
try:
@@ -325,8 +314,8 @@ class SessionPersistenceMixin:
"""Session DB flush, session log and trajectory persistence (see module docstring)."""
def _apply_persist_user_message_override(self, messages: List[Dict]) -> None:
"""Rewrite the current-turn user message in place: some paths send an API-only variant that
must not leak into transcripts or resumed history."""
"""Rewrite the current-turn user message in place: some paths send an API-only variant that must not
leak into transcripts or resumed history."""
idx = getattr(self, "_persist_user_message_idx", None)
override = getattr(self, "_persist_user_message_override", None)
timestamp = getattr(self, "_persist_user_message_timestamp", None)
@@ -345,10 +334,9 @@ class SessionPersistenceMixin:
msg["platform_message_id"] = platform_id
def _persist_session(self, messages: List[Dict], conversation_history: List[Dict] = None):
"""Save session state to both JSON log and SQLite on any exit path. Trailing empty-response
scaffolding is dropped from the live list; the persist override is applied to the DB row only."""
"""Save to JSON log and SQLite on any exit path. Trailing empty-response scaffolding is dropped from
the live list; the persist override is applied to the DB row only."""
from agent.agent_runtime_helpers import note_turn_persisted
with _persist_lock(self):
self._drop_trailing_empty_response_scaffolding(messages)
self._session_messages = messages
@@ -360,9 +348,9 @@ class SessionPersistenceMixin:
note_turn_persisted(self)
def _drop_trailing_empty_response_scaffolding(self, messages: List[Dict]) -> None:
"""Remove empty-response retry scaffolding from the tail, then (only if any was present) rewind
the tool-result / assistant(tool_calls) pair the failed iteration left hanging — otherwise the
next user turn lands as ``...tool, user`` and providers return empty content forever."""
"""Pop empty-response retry scaffolding from the tail, then (only if any was present) rewind the
tool-result / assistant(tool_calls) pair the failed iteration left hanging — otherwise the next user
turn lands as ``...tool, user`` and providers return empty content forever."""
def tail(*keys: str) -> bool:
return bool(messages) and isinstance(messages[-1], dict) and any(messages[-1].get(k) for k in keys)
@@ -391,26 +379,24 @@ class SessionPersistenceMixin:
def _flush_messages_to_session_db_unlocked(
self, messages: List[Dict], conversation_history: Optional[List[Dict]] = None, _adoption_budget: int = 1,
):
"""Persist un-flushed messages to SQLite. Dedup is the intrinsic ``_DB_PERSISTED_MARKER`` on
each written dict — not positional slices (drift after sequence repair) nor an ``id(msg)`` set
(address reuse). The persist override touches the written row only. A compression-closed
session adopts its live tip and retries exactly once."""
# Persistence-isolated agents (background review fork) share the parent's session_id for cache
# warmth; a write here would land the curator's turn in the user's real history.
"""Persist un-flushed messages to SQLite. Dedup is the intrinsic ``_DB_PERSISTED_MARKER`` on each written
dict — not positional slices (drift after sequence repair) nor an ``id(msg)`` set (address reuse). The
persist override touches the written row only. A compression-closed session adopts its live tip and
retries exactly once."""
# Persistence-isolated agents (background review fork) share the parent's session_id for cache warmth;
# a write here would land the curator's turn in the user's real history.
if getattr(self, "_persist_disabled", False) or not self._session_db:
return None
batch_rows: List[Dict[str, Any]] = []
try:
# Retry row creation if the earlier attempt failed transiently.
if not self._session_db_created:
if not self._session_db_created: # retry row creation if the earlier attempt failed transiently
self._ensure_db_session()
batch_rows, batch_msgs = _db_flush_collect(self, messages, conversation_history)
_db_flush_write(self, batch_rows, batch_msgs)
# Markers are now the sole truth; reset the one-shot seed so no id() outlives this flush.
self._flushed_db_message_ids = set()
self._last_flushed_db_idx = len(messages)
# Snapshot for the bounded scan — only on full success, so a partially-processed list can
# never be treated as settled.
# Snapshot for the bounded scan — only on full success, so a partial list is never treated as settled.
self._db_flush_scan_prefix = messages[:]
return True
except Exception as e:
@@ -462,15 +448,14 @@ class SessionPersistenceMixin:
]
def _save_session_log(self, messages: List[Dict[str, Any]] = None):
"""Optional per-session JSON snapshot (``sessions.write_json_snapshots``, default False) for
external tooling; state.db is canonical. Rewrites the full list after every persistence point."""
"""Optional per-session JSON snapshot (``sessions.write_json_snapshots``, default False) for external
tooling; state.db is canonical. Rewrites the full list after every persistence point."""
if not getattr(self, "_session_json_enabled", False):
return
messages = messages or self._session_messages
if not messages:
return
# Re-derive the path each call so /branch and /compress land in the right file.
try:
try: # re-derive the path each call so /branch and /compress land in the right file
log_file = self.logs_dir / f"session_{_safe_session_filename_component(self.session_id)}.json"
except Exception:
return