From 48542259037d8811d8b0cd97254211146eb68da8 Mon Sep 17 00:00:00 2001 From: kshitijk4poor <82637225+kshitijk4poor@users.noreply.github.com> Date: Sun, 27 Sep 2026 13:15:27 +0530 Subject: [PATCH] fix(persistence): version transcript rows by digest, not a row copy The CAS row snapshot was a full copy of each message's durable payload riding on the live dict. The rough token estimator priced it (about 2x estimates -> premature compaction) and it doubled transcript memory. Replace it with a 16-byte blake2b digest of the repair columns. The compare now runs in Python against the target row already read inside the BEGIN IMMEDIATE transaction, followed by a plain UPDATE. Also: - add _db_row_snapshot to PERSISTENCE_ONLY_MESSAGE_FIELDS so the estimator and the outbound request builder both drop it - derive _REPAIR_COLUMNS/_SYNC_FIELDS from _MESSAGE_WRITE_COLUMNS - use hermes_state_common._placeholders - drop the dead resume-path stamp (the SELECT has no token_count, so it was always None) and the dead tool name assignment in _decoded_repair_row - keep the digest out of divert JSONL The kept active-row test now pins estimate stability across a flush and the survival of a concurrent writer's row. It goes red on the old prod files and red when the digest compare is removed. --- agent/message_metadata.py | 4 +- agent/session_persistence.py | 8 +- agent/transcript_repair.py | 78 +++++++++++-------- hermes_state_messages.py | 6 -- .../test_tool_call_incremental_persistence.py | 25 +++++- 5 files changed, 73 insertions(+), 48 deletions(-) diff --git a/agent/message_metadata.py b/agent/message_metadata.py index 687942156a..a6d2adb26a 100644 --- a/agent/message_metadata.py +++ b/agent/message_metadata.py @@ -11,7 +11,9 @@ from typing import Any, MutableMapping, Optional, TypeVar # outgoing copy and the token estimator ignores them: one set, so an estimate # never prices bytes the provider never receives (an edit's inline_diff in # display_metadata is ~9KB and would trigger premature compaction). -PERSISTENCE_ONLY_MESSAGE_FIELDS = frozenset({"timestamp", "display_kind", "display_metadata", "_row_id"}) +PERSISTENCE_ONLY_MESSAGE_FIELDS = frozenset( + {"timestamp", "display_kind", "display_metadata", "_row_id", "_db_row_snapshot"} +) _Message = TypeVar("_Message", bound=MutableMapping[str, Any]) diff --git a/agent/session_persistence.py b/agent/session_persistence.py index e7a356caee..738be15f1a 100644 --- a/agent/session_persistence.py +++ b/agent/session_persistence.py @@ -225,8 +225,8 @@ def _db_flush_row(agent, msg: Dict, is_current_turn_user: bool) -> Dict[str, Any } if isinstance(msg.get("_row_id"), int): row["_row_id"] = msg["_row_id"] - if isinstance(msg.get(_DB_ROW_SNAPSHOT), dict): - row[_DB_ROW_SNAPSHOT] = dict(msg[_DB_ROW_SNAPSHOT]) + if isinstance(msg.get(_DB_ROW_SNAPSHOT), str): + row[_DB_ROW_SNAPSHOT] = msg[_DB_ROW_SNAPSHOT] return row @@ -370,7 +370,9 @@ def _db_flush_failed(agent, e: Exception, batch_rows: List[Dict[str, Any]], adop if isinstance(e, (StateDbReplacedError, StateDbCorruptError)): # A replaced/quarantined handle will not take this batch again — keep it on disk. try: - divert_session_transcript_jsonl(getattr(agent, "session_id", "") or "", batch_rows) + # The CAS digest is local repair bookkeeping, not transcript payload. + divert_session_transcript_jsonl(getattr(agent, "session_id", "") or "", + [{k: v for k, v in r.items() if k != _DB_ROW_SNAPSHOT} for r in batch_rows]) except Exception: logger.warning("JSONL divert failed after state.db %s for %s", agent._last_persistence_error_cause, getattr(agent, "session_id", None), exc_info=True) diff --git a/agent/transcript_repair.py b/agent/transcript_repair.py index 911518080f..6ca91ebc6f 100644 --- a/agent/transcript_repair.py +++ b/agent/transcript_repair.py @@ -4,34 +4,49 @@ clone lookup) and sync markers after commit.""" from __future__ import annotations +import hashlib +import json import sqlite3 from typing import Any, Callable, Dict, List, Mapping, Optional from agent.context_compressor import _DB_PERSISTED_MARKER +from hermes_state_common import _placeholders _DB_ROW_SNAPSHOT = "_db_row_snapshot" _CANONICAL_ROW = "_canonical_row" -_REPAIR_COLUMNS = ( - "content", "tool_call_id", "tool_calls", "tool_name", "effect_disposition", "token_count", - "finish_reason", "reasoning", "reasoning_content", "reasoning_details", "codex_reasoning_items", - "codex_message_items", "platform_message_id", "observed", "_compressed_summary", "api_content", - "display_kind", "display_metadata", -) -_SYNC_FIELDS = ( - "role", "content", "tool_call_id", "tool_calls", "tool_name", "effect_disposition", "token_count", - "finish_reason", "reasoning", "reasoning_content", "reasoning_details", "codex_reasoning_items", - "codex_message_items", "message_id", "platform_message_id", "observed", - "_compressed_summary", "api_content", "display_kind", "display_metadata", -) +def _write_columns() -> tuple: + # Late import: hermes_state_messages imports this module lazily inside its methods. + from hermes_state_messages import _MESSAGE_WRITE_COLUMNS + + return _MESSAGE_WRITE_COLUMNS -def transcript_row_snapshot(row: Mapping[str, Any]) -> Optional[Dict[str, Any]]: - """Serialized durable values used as the CAS version, or ``None`` for a partial SELECT.""" - keys = set(row.keys()) if hasattr(row, "keys") else set(row) +# Durable payload columns a row-addressed rewrite may change: every INSERT column except row identity, +# role, active flag and the ones owned by the display index / timestamp / session linkage. +_NON_PAYLOAD_COLUMNS = frozenset({"session_id", "role", "timestamp", "active", "display_identity"}) +_REPAIR_COLUMNS = tuple(c for c in _write_columns() if c not in _NON_PAYLOAD_COLUMNS) +_SYNC_FIELDS = ("role", "message_id") + _REPAIR_COLUMNS + + +def _canonical_value(value: Any) -> Any: + return value.hex() if isinstance(value, (bytes, bytearray, memoryview)) else value + + +def transcript_row_snapshot(row: Mapping[str, Any]) -> Optional[str]: + """Fixed-size digest of the durable repair columns (the CAS version), or ``None`` for a partial SELECT. + + A digest rather than a row copy: the value rides on live message dicts, so a full copy would double + transcript memory and anything that prices dict bytes. + """ + keys = set(row.keys()) if not set(_REPAIR_COLUMNS) <= keys: return None - return {column: row[column] for column in _REPAIR_COLUMNS} + canonical = json.dumps( + [_canonical_value(row[column]) for column in _REPAIR_COLUMNS], + ensure_ascii=True, separators=(",", ":"), default=str, + ) + return hashlib.blake2b(canonical.encode("utf-8"), digest_size=16).hexdigest() def is_content_blank(content: Any) -> bool: @@ -75,8 +90,10 @@ def resolve_and_repair_transcript_batch( msg["_row_id"] = target_id serialized = serialize_message_fn(msg, float(target_row["timestamp"])) expected = msg.get(_DB_ROW_SNAPSHOT) - has_snapshot = isinstance(expected, dict) and all(column in expected for column in _REPAIR_COLUMNS) - repaired = _compare_and_swap_row(conn, session_id, target_row, serialized, expected) if has_snapshot else False + has_snapshot = isinstance(expected, str) + if has_snapshot and transcript_row_snapshot(target_row) == expected: + # The caller holds BEGIN IMMEDIATE, so the row cannot change between this compare and the UPDATE. + _rewrite_row(conn, session_id, target_row, serialized) if not has_snapshot and role == "assistant" and is_content_blank(decode_content_fn(target_row["content"])): # Blank assistant rows are the pre-existing interrupted-stream repair path. Keep its narrow # content-only CAS for live dicts that predate durable row snapshots. @@ -94,14 +111,13 @@ def resolve_and_repair_transcript_batch( return inserted_rows -def _compare_and_swap_row( +def _rewrite_row( conn: sqlite3.Connection, session_id: str, target_row: Mapping[str, Any], serialized: Mapping[str, Any], - expected: Mapping[str, Any], -) -> bool: - """Rewrite one durable payload only while it still equals the live dict's last committed snapshot.""" +) -> None: + """Rewrite one durable payload whose digest matched the live dict's last committed version.""" old_identity = target_row["display_identity"] old_peer_ids = [ int(row["id"]) @@ -113,17 +129,11 @@ def _compare_and_swap_row( ] if old_identity is not None else [] assignments = ", ".join(f"{column} = ?" for column in _REPAIR_COLUMNS) - predicates = " AND ".join(f"{column} IS ?" for column in _REPAIR_COLUMNS) - params = [serialized[column] for column in _REPAIR_COLUMNS] - params += [int(target_row["id"]), session_id] - params += [expected[column] for column in _REPAIR_COLUMNS] - cur = conn.execute( - f"UPDATE messages SET {assignments} WHERE id = ? AND session_id = ? AND {predicates}", params, + conn.execute( + f"UPDATE messages SET {assignments} WHERE id = ? AND session_id = ?", + [*(serialized[column] for column in _REPAIR_COLUMNS), int(target_row["id"]), session_id], ) - if cur.rowcount != 1: - return False _restore_display_index(conn, session_id, target_row, serialized, old_identity, old_peer_ids) - return True def _restore_display_index( @@ -136,7 +146,7 @@ def _restore_display_index( ) -> None: """Restore display identities/orders invalidated by the payload-update trigger.""" if old_peer_ids: - placeholders = ", ".join("?" for _ in old_peer_ids) + placeholders = _placeholders(old_peer_ids) old_order = min(old_peer_ids) conn.execute( f"UPDATE messages SET display_identity = ?, display_order = ? " @@ -213,8 +223,8 @@ def sync_flushed_message_markers(batch_msgs: List[Dict[str, Any]], batch_rows: L written["_row_id"] = row["_row_id"] if isinstance(row.get("timestamp"), (int, float)): written["timestamp"] = row["timestamp"] - if isinstance(row.get(_DB_ROW_SNAPSHOT), dict): - written[_DB_ROW_SNAPSHOT] = dict(row[_DB_ROW_SNAPSHOT]) + if isinstance(row.get(_DB_ROW_SNAPSHOT), str): + written[_DB_ROW_SNAPSHOT] = row[_DB_ROW_SNAPSHOT] canonical = row.get(_CANONICAL_ROW) if isinstance(canonical, dict): for key in _SYNC_FIELDS: diff --git a/hermes_state_messages.py b/hermes_state_messages.py index 22dd95337d..13fb5759e2 100644 --- a/hermes_state_messages.py +++ b/hermes_state_messages.py @@ -317,8 +317,6 @@ class SessionMessagesMixin: if row["platform_message_id"] is not None: msg["message_id"] = row["platform_message_id"] msg["platform_message_id"] = row["platform_message_id"] - if row["role"] == "tool" and row["tool_name"] is not None: - msg["name"] = row["tool_name"] if row["observed"]: msg["observed"] = True if row["_compressed_summary"]: @@ -1499,11 +1497,7 @@ class SessionMessagesMixin: # durable-snapshot adoption, incremental persists with no history arg) would otherwise re-append # the ENTIRE transcript on flush. if include_row_ids and row["id"] is not None: - from agent.transcript_repair import _DB_ROW_SNAPSHOT, transcript_row_snapshot - msg["_row_id"] = row["id"] - if (snapshot := transcript_row_snapshot(row)) is not None: - msg[_DB_ROW_SNAPSHOT] = snapshot msg.update((col, row[col]) for col in ("api_content", "display_kind") if row[col]) if row["display_metadata"] and (decoded := self._decode_display_metadata(row["display_metadata"])) is not None: msg["display_metadata"] = decoded diff --git a/tests/agent/test_tool_call_incremental_persistence.py b/tests/agent/test_tool_call_incremental_persistence.py index 402e02c15f..cbd3789f95 100644 --- a/tests/agent/test_tool_call_incremental_persistence.py +++ b/tests/agent/test_tool_call_incremental_persistence.py @@ -22,6 +22,7 @@ read snapshots captured at flush time, so removing any production flush call makes the corresponding assertion fail. """ +import sqlite3 import copy from types import SimpleNamespace from pathlib import Path @@ -749,20 +750,33 @@ def test_flush_stale_row_id_from_other_session_does_not_fill_child_blank(tmp_pat def test_flush_sanitized_active_user_and_tool_rows_do_not_append_duplicates(tmp_path): - """The outbound sanitizer pops ``_db_persisted``; active rows keep their ``_row_id`` and are not re-inserted.""" + """The outbound sanitizer pops ``_db_persisted``; active rows keep their ``_row_id`` and are not re-inserted. + + The row-version digest must not be priced by the token estimator, and a row another writer changed after our + flush must survive the re-flush (the concurrent winner is adopted, never clobbered or duplicated). + """ from agent.message_sanitization import _sanitize_messages_surrogates + from agent.model_metadata import estimate_messages_tokens_rough agent = _make_agent() session_id = "sess-sanitized-active-rows" db = _attach_real_session_db(agent, tmp_path / "state.db", session_id) messages = [ - {"role": "user", "content": "hi \ud800 there"}, + {"role": "user", "content": "hi \ud800 there " + "x" * 4000}, {"role": "assistant", "content": "ok"}, {"role": "tool", "tool_call_id": "c1", "name": "terminal", "content": "r \ud800"}, ] + estimate_before = estimate_messages_tokens_rough(messages) agent._flush_messages_to_session_db(messages) + assert estimate_messages_tokens_rough(messages) - estimate_before < 50 durable_ids = [message["_row_id"] for message in messages] + # Another writer replaces the tool row after our flush; our live dict still carries the old version. + other = sqlite3.connect(tmp_path / "state.db") + with other: + other.execute("UPDATE messages SET content = ? WHERE id = ?", ("winner", durable_ids[2])) + other.close() + assert _sanitize_messages_surrogates(messages) is True assert not any(message.get("_db_persisted") for message in (messages[0], messages[2])) agent._db_flush_scan_prefix = None @@ -772,6 +786,9 @@ def test_flush_sanitized_active_user_and_tool_rows_do_not_append_duplicates(tmp_ rows = db.get_messages(session_id, include_inactive=True) assert [row["id"] for row in rows] == durable_ids assert [message["_row_id"] for message in messages] == durable_ids + assert rows[0]["content"].startswith("hi \ufffd there") + assert rows[2]["content"] == "winner" + assert messages[2]["content"] == "winner" def test_flush_sanitized_archived_user_and_tool_rows_do_not_append_duplicates(tmp_path): @@ -821,8 +838,8 @@ def test_flush_sanitized_archived_user_and_tool_rows_do_not_append_duplicates(tm assert len(matching) == 1 assert matching[0]["id"] == durable_ids[role] assert matching[0]["timestamp"] == durable_timestamps[role] - assert matching[0]["active"] in (0, False) - assert matching[0]["compacted"] in (1, True) + assert matching[0]["active"] == 0 + assert matching[0]["compacted"] == 1 live = next(message for message in messages if message["role"] == role) assert live["_row_id"] == durable_ids[role] assert live["_db_persisted"] is True