diff --git a/hermes_state_compression.py b/hermes_state_compression.py index eb9dfb5f49..1110b66880 100644 --- a/hermes_state_compression.py +++ b/hermes_state_compression.py @@ -9,7 +9,7 @@ import json import logging import sqlite3 import time -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Tuple from hermes_state_common import _sql_session_last_active, is_automatic_end_reason @@ -27,6 +27,30 @@ def _ended_by_compression(row) -> bool: return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" +def _cooldown_row(exists: bool, cooldown_until, error) -> Dict[str, Any]: + return { + "session_exists": exists, + "cooldown_until": float(cooldown_until) if cooldown_until is not None else None, + "error": error} + + +def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float, + stale) -> Tuple[bool, Optional[str]]: + """Single-transaction lease claim: DELETE a stale holder's row (``stale(holder, + expires_at)``), INSERT OR IGNORE ours, then SELECT to confirm ownership (INSERT OR + IGNORE gives no rowcount signal). Returns ``(acquired, reclaimed_holder)``.""" + reclaimed_holder = None + row = conn.execute(f"SELECT holder, expires_at FROM {table} WHERE {key_col} = ?", (key,)).fetchone() + if row is not None and stale(row["holder"], row["expires_at"]): + conn.execute(f"DELETE FROM {table} WHERE {key_col} = ? AND holder = ?", (key, row["holder"])) + reclaimed_holder = row["holder"] + conn.execute( + f"INSERT OR IGNORE INTO {table} ({key_col}, holder, acquired_at, expires_at) VALUES (?, ?, ?, ?)", + (key, holder, now, expires_at)) + owner = conn.execute(f"SELECT holder FROM {table} WHERE {key_col} = ?", (key,)).fetchone() + return owner is not None and owner["holder"] == holder, reclaimed_holder + + class SessionCompressionMixin: """Compression lineage, cooldown/streak counters, locks and turn leases.""" @@ -96,16 +120,13 @@ class SessionCompressionMixin: deleted = conn.execute( "DELETE FROM compression_locks " "WHERE session_id = ? AND holder = ? AND expires_at = ?", - (session_id, lock_row["holder"], expires_at), - ) + (session_id, lock_row["holder"], expires_at)) if deleted.rowcount != 1: return False updated = conn.execute( "UPDATE sessions SET ended_at = NULL, end_reason = NULL " - "WHERE id = ? AND ended_at IS NOT NULL " - "AND end_reason = 'compression'", - (session_id,), - ) + "WHERE id = ? AND ended_at IS NOT NULL AND end_reason = 'compression'", + (session_id,)) # rowcount==1 is guaranteed by the parent SELECT in this same txn. A False # return added past this point must raise instead: the lease DELETE above # commits unless _do raises. @@ -115,7 +136,10 @@ class SessionCompressionMixin: def _publish_child_session_row(self, conn, parent, *, parent_session_id, child_session_id, source, model, model_config, system_prompt, cwd, profile_name) -> None: - """INSERT the compression child's ``sessions`` row copied from *parent*.""" + """INSERT the compression child's ``sessions`` row copied from *parent*. Same contract + as _insert_session_row's compression-fork backfill: the child stays on the parent's + profile and keeps gateway routing/origin columns; no owner on either side -> this + store's profile.""" system_prompt_hash = self._store_system_prompt(conn, system_prompt) conn.execute( """INSERT INTO sessions ( @@ -126,28 +150,12 @@ class SessionCompressionMixin: thread_id, display_name, origin_json, started_at ) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", ( - child_session_id, - source, - model, - json.dumps(model_config) if model_config else None, - system_prompt_hash, - parent_session_id, - cwd or parent["cwd"], - parent["git_branch"], + child_session_id, source, model, json.dumps(model_config) if model_config else None, + system_prompt_hash, parent_session_id, cwd or parent["cwd"], parent["git_branch"], parent["git_repo_root"], - # Same contract as _insert_session_row's compression-fork backfill: the - # child stays on the parent's profile and keeps gateway routing/origin - # columns; no owner on either side -> this store's profile. profile_name or parent["profile_name"] or self._own_profile_name(), - parent["user_id"], - parent["session_key"], - parent["chat_id"], - parent["chat_type"], - parent["thread_id"], - parent["display_name"], - parent["origin_json"], - time.time(), - ), + parent["user_id"], parent["session_key"], parent["chat_id"], parent["chat_type"], + parent["thread_id"], parent["display_name"], parent["origin_json"], time.time()), ) def publish_compression_child( @@ -180,8 +188,7 @@ class SessionCompressionMixin: conn.execute( "UPDATE compression_locks SET expires_at = ? " "WHERE session_id = ? AND holder = ?", - (time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder), - ) + (time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder)) lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone() if require_compression_lease and ( lock_row is None @@ -209,10 +216,8 @@ class SessionCompressionMixin: # Deliberate boundaries still fail closed. if is_automatic_end_reason(parent["end_reason"]): conn.execute( - "UPDATE sessions SET ended_at = NULL, end_reason = NULL " - "WHERE id = ?", - (parent_session_id,), - ) + "UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?", + (parent_session_id,)) else: raise RuntimeError(f"Compression parent already ended: {parent_session_id}") if not messages: @@ -220,23 +225,17 @@ class SessionCompressionMixin: self._publish_child_session_row( conn, parent, parent_session_id=parent_session_id, child_session_id=child_session_id, source=source, model=model, model_config=model_config, system_prompt=system_prompt, - cwd=cwd, profile_name=profile_name, - ) + cwd=cwd, profile_name=profile_name) total_messages, total_tool_calls = self._insert_message_rows(conn, child_session_id, messages) if watermark is not None: # Clone the parent's concurrent tail into the child after the handoff; # originals stay in the closed parent for lineage recovery. - _ceiling_clause = "" - _params: list = [parent_session_id, int(watermark)] - if watermark_ceiling is not None: - _ceiling_clause = " AND id <= ?" - _params.append(int(watermark_ceiling)) + bounded = watermark_ceiling is not None tail_ids, tail_tool_calls = self._tail_rows_after_watermark( - conn, - "SELECT id, tool_calls FROM messages " + conn, "SELECT id, tool_calls FROM messages " "WHERE session_id = ? AND active = 1 AND id > ?" - f"{_ceiling_clause} ORDER BY id", - _params, + f"{' AND id <= ?' if bounded else ''} ORDER BY id", + [parent_session_id, int(watermark), *([int(watermark_ceiling)] if bounded else [])], ) if tail_ids: self._clone_message_rows(conn, tail_ids, session_id=child_session_id) @@ -244,13 +243,10 @@ class SessionCompressionMixin: total_tool_calls += tail_tool_calls conn.execute( "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", - (total_messages, total_tool_calls, child_session_id), - ) + (total_messages, total_tool_calls, child_session_id)) updated = conn.execute( "UPDATE sessions SET ended_at = ?, end_reason = 'compression' " - "WHERE id = ? AND ended_at IS NULL", - (time.time(), parent_session_id), - ) + "WHERE id = ? AND ended_at IS NULL", (time.time(), parent_session_id)) if updated.rowcount != 1: raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}") @@ -276,8 +272,7 @@ class SessionCompressionMixin: " AND compression_failure_cooldown_until > ? " "THEN compression_failure_cooldown_until ELSE ? END, " "compression_failure_error = ? WHERE id = ?", - (cooldown_until, cooldown_until, error, session_id), - ) + (cooldown_until, cooldown_until, error, session_id)) def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]: """Return the active (unexpired) compression-failure cooldown, or None.""" @@ -285,24 +280,17 @@ class SessionCompressionMixin: return None now = time.time() row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) - if row is None or row[0] is None: + if row is None or row[0] is None or float(row[0]) <= now: return None - cooldown_until = float(row[0]) - if cooldown_until <= now: - return None - return {"cooldown_until": cooldown_until, "remaining_seconds": cooldown_until - now, "error": row[1]} + return {"cooldown_until": float(row[0]), "remaining_seconds": float(row[0]) - now, "error": row[1]} def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]: """Exact stored cooldown columns, no expiry filtering, so compression cancellation can roll back an expired, partially-null, or absent row exactly.""" row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None if row is None: - return {"session_exists": False, "cooldown_until": None, "error": None} - return { - "session_exists": True, - "cooldown_until": float(row[0]) if row[0] is not None else None, - "error": row[1], - } + return _cooldown_row(False, None, None) + return _cooldown_row(True, row[0], row[1]) def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None: """Restore and verify an exact cooldown-row snapshot. Unlike record/clear this @@ -318,19 +306,12 @@ class SessionCompressionMixin: def _do(conn): cursor = conn.execute( "UPDATE sessions SET compression_failure_cooldown_until = ?, " - "compression_failure_error = ? WHERE id = ?", - (deadline, error, session_id), - ) + "compression_failure_error = ? WHERE id = ?", (deadline, error, session_id)) if cursor.rowcount != 1: raise RuntimeError(f"compression cooldown rollback session missing: {session_id}") - self._execute_write(_do) actual = self.get_compression_failure_cooldown_row(session_id) - expected = { - "session_exists": True, - "cooldown_until": float(deadline) if deadline is not None else None, - "error": error, - } + expected = _cooldown_row(True, deadline, error) if actual != expected: raise RuntimeError( f"compression cooldown rollback verification failed: " @@ -344,9 +325,7 @@ class SessionCompressionMixin: self._write_sql_logged( "clear_compression_failure_cooldown", session_id, "UPDATE sessions SET compression_failure_cooldown_until = NULL, " - "compression_failure_error = NULL WHERE id = ?", - (session_id,), - ) + "compression_failure_error = NULL WHERE id = ?", (session_id,)) def _read_session_number(self, column: str, session_id: str, cast: type, zero: Any) -> Any: """Read one numeric ``sessions`` column clamped at ``zero``; a missing session, @@ -370,8 +349,7 @@ class SessionCompressionMixin: if session_id: self._write_sql( "UPDATE sessions SET compression_fallback_streak = ? WHERE id = ?", - (max(0, int(streak)), session_id), - ) + (max(0, int(streak)), session_id)) def get_compression_ineffective_count(self, session_id: str) -> int: """Persisted ineffective-compaction strike count — the durable half of the @@ -384,8 +362,7 @@ class SessionCompressionMixin: if session_id: self._write_sql( "UPDATE sessions SET compression_ineffective_count = ? WHERE id = ?", - (max(0, int(count)), session_id), - ) + (max(0, int(count)), session_id)) def get_compression_recovery_deadline(self, session_id: str) -> float: """Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable @@ -402,8 +379,7 @@ class SessionCompressionMixin: normalized = 0.0 self._write_sql( "UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?", - (normalized if normalized > 0.0 else None, session_id), - ) + (normalized or None, session_id)) def refresh_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool: """Extend the compression lock lease if ``holder`` still owns it. @@ -419,8 +395,7 @@ class SessionCompressionMixin: expires_at = time.time() + ttl_seconds try: return self._write_rowcount( - "UPDATE compression_locks SET expires_at = ? " - "WHERE session_id = ? AND holder = ?", + "UPDATE compression_locks SET expires_at = ? WHERE session_id = ? AND holder = ?", (expires_at, session_id, holder), ) > 0 except sqlite3.Error as exc: @@ -442,32 +417,14 @@ class SessionCompressionMixin: expires_at = now + ttl_seconds def _do(conn): - reclaimed_holder = None - row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone() - if row is not None: - current_holder, current_expires_at = row[0], row[1] - if current_expires_at < now or _compression_lock_holder_process_is_dead(current_holder): - conn.execute( - "DELETE FROM compression_locks " - "WHERE session_id = ? AND holder = ?", - (session_id, current_holder), - ) - reclaimed_holder = current_holder - conn.execute( - "INSERT OR IGNORE INTO compression_locks " - "(session_id, holder, acquired_at, expires_at) " - "VALUES (?, ?, ?, ?)", - (session_id, holder, now, expires_at), - ) - row = conn.execute("SELECT holder FROM compression_locks WHERE session_id = ?", (session_id,)).fetchone() - return row is not None and row[0] == holder, reclaimed_holder + return _claim_lease_row( + conn, "compression_locks", "session_id", session_id, holder, now, expires_at, + lambda h, e: e < now or _compression_lock_holder_process_is_dead(h)) try: acquired, reclaimed_holder = self._execute_write(_do) if reclaimed_holder: - logger.warning( - "Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder, - ) + logger.warning("Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder) return bool(acquired) except sqlite3.Error as exc: # False makes the caller skip compression — safe when the lock subsystem is broken. @@ -480,10 +437,8 @@ class SessionCompressionMixin: return self._write_sql_logged( "release_compression_lock", session_id, - "DELETE FROM compression_locks " - "WHERE session_id = ? AND holder = ?", - (session_id, holder), - ) + "DELETE FROM compression_locks WHERE session_id = ? AND holder = ?", + (session_id, holder)) def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str: """Walk compression parents on ``conn`` to the conversation lease key. @@ -496,9 +451,7 @@ class SessionCompressionMixin: def _row(sid: str): row = conn.execute( - "SELECT id, parent_session_id, source, model_config, end_reason " - "FROM sessions WHERE id = ?", - (sid,), + "SELECT id, parent_session_id, source, model_config, end_reason FROM sessions WHERE id = ?", (sid,), ).fetchone() return dict(row) if row else None @@ -537,29 +490,10 @@ class SessionCompressionMixin: def _do(conn): conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) - row = conn.execute( - "SELECT holder, expires_at FROM session_turn_leases " - "WHERE conversation_id = ?", - (conversation_id,), - ).fetchone() - if row is not None: - current_holder = row["holder"] - if float(row["expires_at"]) <= now or _compression_lock_holder_process_is_dead(current_holder): - conn.execute( - "DELETE FROM session_turn_leases " - "WHERE conversation_id = ? AND holder = ?", - (conversation_id, current_holder), - ) - conn.execute( - "INSERT OR IGNORE INTO session_turn_leases " - "(conversation_id, holder, acquired_at, expires_at) " - "VALUES (?, ?, ?, ?)", - (conversation_id, holder, now, expires_at), - ) - owner = conn.execute( - "SELECT holder FROM session_turn_leases WHERE conversation_id = ?", (conversation_id,), - ).fetchone() - return owner is not None and owner["holder"] == holder + return _claim_lease_row( + conn, "session_turn_leases", "conversation_id", conversation_id, holder, now, expires_at, + lambda h, e: float(e) <= now or _compression_lock_holder_process_is_dead(h), + )[0] return bool(self._execute_write(_do, patience_s=patience_s)) @@ -588,8 +522,7 @@ class SessionCompressionMixin: logger.debug("session turn lease should_abort callback failed", exc_info=True) try: if self.try_acquire_session_turn_lease( - session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s, - ): + session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s): return True except sqlite3.Error as exc: # Long holder transactions can exhaust one write-patience budget; keep @@ -620,12 +553,10 @@ class SessionCompressionMixin: def _do(conn): conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) - cursor = conn.execute( + return conn.execute( "UPDATE session_turn_leases SET expires_at = ? " - "WHERE conversation_id = ? AND holder = ?", - (expires_at, conversation_id, holder), - ) - return cursor.rowcount > 0 + "WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder), + ).rowcount > 0 return bool(self._execute_write(_do)) @@ -637,10 +568,8 @@ class SessionCompressionMixin: def _do(conn): conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) conn.execute( - "DELETE FROM session_turn_leases " - "WHERE conversation_id = ? AND holder = ?", - (conversation_id, holder), - ) + "DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?", + (conversation_id, holder)) self._execute_write(_do) @@ -649,10 +578,8 @@ class SessionCompressionMixin: if not session_id: return None row = self._read_one( - "SELECT holder FROM compression_locks " - "WHERE session_id = ? AND expires_at >= ?", - (session_id, time.time()), - ) + "SELECT holder FROM compression_locks WHERE session_id = ? AND expires_at >= ?", + (session_id, time.time())) return None if row is None else row[0] def finalize_orphaned_compression_sessions(self) -> int: @@ -660,10 +587,7 @@ class SessionCompressionMixin: has messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as ``orphaned_compression``. Non-destructive.""" cutoff = time.time() - 604800 # 7 days - - def _do(conn): - now = time.time() - result = conn.execute( + return self._write_rowcount( """ UPDATE sessions SET ended_at = ?, @@ -684,11 +608,8 @@ class SessionCompressionMixin: WHERE m.session_id = sessions.id ) """, - (now, cutoff), - ) - return result.rowcount - - return self._execute_write(_do) or 0 + (time.time(), cutoff), + ) or 0 def get_compression_chain(self, session_id: str) -> List[str]: """Walk the compression-continuation chain forward: root-first through the tip @@ -703,7 +624,7 @@ class SessionCompressionMixin: are still live over stale closed siblings such as ``ws_orphan_reap``.""" current = session_id chain = [current] if current else [] - seen = {current} if current else set() + seen = set(chain) for _ in range(100): # defensive bound; chains this deep are pathological with self._read_ctx() as conn: row = conn.execute( @@ -729,9 +650,7 @@ class SessionCompressionMixin: """, (current,), ).fetchone() - if row is None: - return chain - child_id = row["id"] + child_id = row["id"] if row is not None else None if not child_id or child_id in seen: return chain seen.add(child_id) diff --git a/hermes_state_messages.py b/hermes_state_messages.py index 2a3f57a1f7..71e37665e3 100644 --- a/hermes_state_messages.py +++ b/hermes_state_messages.py @@ -48,15 +48,18 @@ def _json_or(raw: Any, fallback: Any, warning: str) -> Any: return fallback -def _tool_calls_len(raw: Any) -> int: - """Tool-call count of a stored ``tool_calls`` column (0 unless a JSON list).""" +def _tool_calls_len(raw: Any, scalar: int = 0) -> int: + """Tool-call count of a stored ``tool_calls`` column: list length, else *scalar* + for a truthy non-list value, 0 for empty/undecodable.""" if not raw: return 0 try: parsed = json.loads(raw) if isinstance(raw, str) else raw except (TypeError, ValueError): return 0 - return len(parsed) if isinstance(parsed, list) else 0 + if isinstance(parsed, list): + return len(parsed) + return scalar if parsed else 0 def _coerce_timestamp(value: Any, default: float) -> float: @@ -91,18 +94,28 @@ def _ended_by_compression(row) -> bool: return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" +# _rows_to_conversation copies these row columns verbatim when truthy, in this order. +# ``api_content`` is returned VERBATIM (no sanitize/strip): the replay path substitutes +# it to keep the provider prompt cache byte-stable. +_VERBATIM_COLS = ("api_content", "display_kind") +_META_COLS = ("timestamp", "tool_call_id", "tool_name", "effect_disposition") +_ASSISTANT_JSON_COLS = ("reasoning_details", "codex_reasoning_items", "codex_message_items") + + +def _stale_holder(row, now: float) -> bool: + """A lock/lease row whose holder is expired or a provably dead local process.""" + from hermes_state import _compression_lock_holder_process_is_dead + return float(row["expires_at"]) <= now or _compression_lock_holder_process_is_dead(row["holder"]) + + class SessionMessagesMixin: """Message append/replace/rewind, reactions, resume conversations, replay dedupe.""" def _bump_conversation_generation(self, conn, session_id: str, end_reason: str) -> None: - """Advance this peer's conversation generation past a boundary, inside the - transaction that writes the boundary. - - Only ``_RESET_END_REASONS`` count (``compression`` continues one conversation). - The counter never reads the session rows: an aggregate over them could re-emit a - pair once ``delete_session()``/pruning removes an ended row and hand a new - conversation a retired affinity identity. It only ever increments. - """ + """Advance this peer's conversation generation past a boundary, inside the txn that + writes it. Only ``_RESET_END_REASONS`` count (``compression`` continues one + conversation). Never derived from session rows (deletes/prunes could re-emit a + retired affinity identity); it only ever increments.""" if end_reason not in _RESET_END_REASONS: return row = conn.execute("SELECT source, session_key FROM sessions WHERE id = ?", (session_id,)).fetchone() @@ -126,8 +139,7 @@ class SessionMessagesMixin: def _encode_content(cls, content: Any) -> Any: """Serialize list/dict content (multimodal parts) as a sentinel-prefixed JSON string; sqlite3 can only bind str/bytes/int/float/None. Lone surrogates are - scrubbed from text so persistence never fails. Paired with :meth:`_decode_content`. - """ + scrubbed from text so persistence never fails. Paired with :meth:`_decode_content`.""" if isinstance(content, str): # Lone UTF-16 surrogates arrive in tool results scraped from the web. The # upstream sanitizer only cleans the api_messages copy and the recovery @@ -150,8 +162,7 @@ class SessionMessagesMixin: if isinstance(content, str) and content.startswith(cls._CONTENT_JSON_PREFIX): return _json_or( content[len(cls._CONTENT_JSON_PREFIX):], content, - "Failed to decode JSON-encoded message content; returning raw string", - ) + "Failed to decode JSON-encoded message content; returning raw string") return content @staticmethod @@ -162,20 +173,19 @@ class SessionMessagesMixin: return None if isinstance(display_metadata, str): try: - parsed = json.loads(display_metadata) + display_metadata = json.loads(display_metadata) except (json.JSONDecodeError, TypeError): logger.warning("Ignoring non-JSON display metadata on write") return None - if not isinstance(parsed, dict): + if not isinstance(display_metadata, dict): logger.warning("Ignoring non-object display metadata on write") return None - return json.dumps(parsed) - if isinstance(display_metadata, dict): - return json.dumps(display_metadata) - logger.warning( - "Ignoring unexpected display metadata type on write: %s", type(display_metadata).__name__, - ) - return None + elif not isinstance(display_metadata, dict): + logger.warning( + "Ignoring unexpected display metadata type on write: %s", type(display_metadata).__name__, + ) + return None + return json.dumps(display_metadata) @staticmethod def _decode_display_metadata(raw: Any) -> Optional[Dict[str, Any]]: @@ -212,27 +222,19 @@ class SessionMessagesMixin: reject_active_turn_lease: bool = False, reject_active_compression_lock: bool = False, allow_closed_compression_parent: bool = False, ) -> None: - """Transcript-write admission checks, run INSIDE the write txn. - - Shared by every transcript writer so they cannot diverge. Ordinary appends do - NOT check compression_locks: the lock only stops two COMPRESSIONS colliding, and - archive_and_compact() commits against a watermark and clones later rows, so - concurrent appends are safe by construction (blocking them killed turns while a - slow summary held the lease). Destructive user mutations opt in via - ``reject_active_compression_lock`` / ``reject_active_turn_lease`` so a compressor - that captured its watermark cannot resurrect the removed turn. - """ - from hermes_state import CompressionSessionClosedError, SessionCompressionInProgressError, SessionTurnLeaseLostError, _compression_lock_holder_process_is_dead + """Transcript-write admission checks, run INSIDE the write txn (shared by every + writer). Ordinary appends do NOT check compression_locks: the lock only stops two + COMPRESSIONS colliding and archive_and_compact() commits against a watermark, so + concurrent appends are safe (blocking them killed turns during slow summaries). + Destructive user mutations opt in via ``reject_active_*`` so a compressor that + captured its watermark cannot resurrect the removed turn.""" + from hermes_state import CompressionSessionClosedError, SessionCompressionInProgressError, SessionTurnLeaseLostError if reject_active_compression_lock: active_lock = conn.execute(_COMPRESSION_LOCK_ROW_SQL, (session_id,)).fetchone() if active_lock is not None: - current_holder = active_lock["holder"] - if ( - float(active_lock["expires_at"]) <= time.time() - or _compression_lock_holder_process_is_dead(current_holder) - ): - conn.execute(_DELETE_COMPRESSION_LOCK_SQL, (session_id, current_holder)) - elif current_holder != compression_lock_holder: + if _stale_holder(active_lock, time.time()): + conn.execute(_DELETE_COMPRESSION_LOCK_SQL, (session_id, active_lock["holder"])) + elif active_lock["holder"] != compression_lock_holder: raise SessionCompressionInProgressError( f"Session {session_id!r} is being compressed by another writer" ) @@ -255,21 +257,15 @@ class SessionMessagesMixin: (now + max(0.1, float(turn_lease_ttl_seconds)), conversation_id, turn_lease_holder), ) elif lease is not None: - current_holder = lease["holder"] - if ( - float(lease["expires_at"]) <= now - or _compression_lock_holder_process_is_dead(current_holder) - ): - # Same reclaim rule as acquisition (expired or provably dead owner); - # deleting here also fences a stale late flush after the mutation. - conn.execute( - "DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?", - (conversation_id, current_holder), - ) - else: + if not _stale_holder(lease, now): raise SessionTurnLeaseLostError( f"Session has an active turn lease; refusing transcript mutation for {session_id!r}" ) + # Same reclaim rule as acquisition (expired or provably dead owner); + # deleting here also fences a stale late flush after the mutation. + conn.execute( + "DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?", + (conversation_id, lease["holder"])) session = conn.execute(_ENDED_BY_COMPRESSION_SQL, (session_id,)).fetchone() if _ended_by_compression(session) and not allow_closed_compression_parent: raise CompressionSessionClosedError(session_id) @@ -281,16 +277,10 @@ class SessionMessagesMixin: """Bind values for ``_INSERT_MESSAGE_SQL`` from one message dict. *tool_calls* is the already-parsed value (see ``_parse_tool_calls``). - *keep_reasoning* False stores NULL for every reasoning/codex column. - """ + *keep_reasoning* False stores NULL for every reasoning/codex column.""" from hermes_state import _scrub_surrogates - - def _str_or_none(value): - return _scrub_surrogates(value) if isinstance(value, str) else None - - def _reasoning(key): - return msg.get(key) if keep_reasoning else None - + _str_or_none = lambda v: _scrub_surrogates(v) if isinstance(v, str) else None # noqa: E731 + _reasoning = lambda key: msg.get(key) if keep_reasoning else None # noqa: E731 return ( session_id, role, @@ -329,39 +319,23 @@ class SessionMessagesMixin: compression_lock_holder: Optional[str] = None, turn_lease_holder: Optional[str] = None, turn_lease_ttl_seconds: float = 300.0, ) -> int: - """Append one message; returns the row id. Bumps ``message_count`` (and - ``tool_call_count`` when tool_calls are present). - - ``platform_message_id`` is the platform's own id (Telegram update_id, Yuanbao - msg_id) used by recall-style flows. ``api_content`` is the byte-fidelity - sidecar — the exact string sent to the API when it differed from ``content`` — - stored as sent except lone surrogates (which the loop scrubs anyway). - """ - msg = { - "content": content, "tool_name": tool_name, "tool_call_id": tool_call_id, - "token_count": token_count, "finish_reason": finish_reason, "reasoning": reasoning, - "reasoning_content": reasoning_content, "reasoning_details": reasoning_details, - "codex_reasoning_items": codex_reasoning_items, "codex_message_items": codex_message_items, - "platform_message_id": platform_message_id, "message_id": platform_message_id, - "observed": observed, "effect_disposition": effect_disposition, - "_compressed_summary": _compressed_summary, "api_content": api_content, - "display_kind": display_kind, "display_metadata": display_metadata, - } + """Append one message; returns the row id and bumps the session counters. + ``platform_message_id``: the platform's own id (recall-style flows). + ``api_content``: byte-fidelity sidecar — the exact string sent to the API when it + differed from ``content`` — stored as sent except lone surrogates.""" + msg = dict(locals()) # every keyword above is a message-dict field of the same name # Encode outside the write txn (display metadata first: log-order parity). - display_metadata_json = self._encode_display_metadata(display_metadata) - msg["display_metadata"] = display_metadata_json + msg["display_metadata"] = self._encode_display_metadata(display_metadata) tool_calls = _parse_tool_calls(tool_calls) message_timestamp = _coerce_timestamp(timestamp, time.time()) num_tool_calls = _tool_calls_count(tool_calls) params = self._message_row_params( - session_id, role, msg, tool_calls, message_timestamp, keep_reasoning=True, - ) + session_id, role, msg, tool_calls, message_timestamp, keep_reasoning=True) def _do(conn): self._check_transcript_write_guards( conn, session_id, compression_lock_holder, - turn_lease_holder=turn_lease_holder, turn_lease_ttl_seconds=turn_lease_ttl_seconds, - ) + turn_lease_holder=turn_lease_holder, turn_lease_ttl_seconds=turn_lease_ttl_seconds) msg_id = conn.execute(_INSERT_MESSAGE_SQL, params).lastrowid if num_tool_calls > 0: conn.execute( @@ -384,13 +358,10 @@ class SessionMessagesMixin: compression_lock_holder: Optional[str] = None, turn_lease_holder: Optional[str] = None, chunk_rows: Optional[int] = None, turn_lease_ttl_seconds: float = 300.0, ) -> int: - """Append *messages* (``_insert_message_rows`` dict shape) in ONE write txn. - - All rows land or none do; the admission guards run once for the batch. - ``chunk_rows`` bounds transaction size for LARGE copies (branch seeds: FTS - triggers run per row, 10k rows ≈ 2.4s under one lock) — commits in chunks with - the old per-row-loop recovery semantics. Returns the inserted row count. - """ + """Append *messages* (``_insert_message_rows`` dict shape) in ONE write txn: all + rows land or none, guards run once. ``chunk_rows`` bounds txn size for LARGE + copies (branch seeds; FTS triggers run per row) by committing in chunks. + Returns the inserted row count.""" if not messages: return 0 if chunk_rows is not None and len(messages) > chunk_rows: @@ -399,25 +370,21 @@ class SessionMessagesMixin: session_id, messages[start:start + chunk_rows], compression_lock_holder=compression_lock_holder, turn_lease_holder=turn_lease_holder, - turn_lease_ttl_seconds=turn_lease_ttl_seconds, - ) + turn_lease_ttl_seconds=turn_lease_ttl_seconds) for start in range(0, len(messages), chunk_rows) ) def _do(conn): self._check_transcript_write_guards( conn, session_id, compression_lock_holder, - turn_lease_holder=turn_lease_holder, turn_lease_ttl_seconds=turn_lease_ttl_seconds, - ) + turn_lease_holder=turn_lease_holder, turn_lease_ttl_seconds=turn_lease_ttl_seconds) from agent.transcript_repair import resolve_and_repair_transcript_batch - inserted_rows = resolve_and_repair_transcript_batch( conn, session_id, messages, - encode_content_fn=self._encode_content, decode_content_fn=self._decode_content, + encode_content_fn=self._encode_content, decode_content_fn=self._decode_content) + inserted, tool_calls_total = ( + self._insert_message_rows(conn, session_id, inserted_rows) if inserted_rows else (0, 0) ) - inserted = tool_calls_total = 0 - if inserted_rows: - inserted, tool_calls_total = self._insert_message_rows(conn, session_id, inserted_rows) if tool_calls_total > 0: conn.execute( """UPDATE sessions SET message_count = message_count + ?, @@ -427,8 +394,7 @@ class SessionMessagesMixin: elif inserted > 0: conn.execute( "UPDATE sessions SET message_count = message_count + ? WHERE id = ?", - (inserted, session_id), - ) + (inserted, session_id)) return inserted return self._execute_write(_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S) @@ -470,12 +436,9 @@ class SessionMessagesMixin: def set_message_reaction( self, session_id: str, message_row_id: int, emoji: Optional[str], *, author: str = "user", ) -> Optional[List[Dict[str, Any]]]: - """Set (or with ``emoji=None`` clear) *author*'s reaction on one message. - - Tapback semantics: one reaction per author per message; the same emoji again - clears it, a different one replaces it. Returns the message's reaction list - after the write, or ``None`` when the row isn't part of *session_id*. - """ + """Set (``emoji=None``: clear) *author*'s reaction. Tapback semantics: one per + author per message; the same emoji again clears it, a different one replaces it. + Returns the reaction list after the write, or ``None`` for a foreign row.""" from hermes_state import _scrub_surrogates if not session_id or message_row_id is None: return None @@ -491,8 +454,7 @@ class SessionMessagesMixin: existing = self._reaction_list(meta) reactions = [r for r in existing if r.get("author") != author] previous = next((r for r in existing if r.get("author") == author), None) - toggling_off = emoji is not None and previous is not None and previous.get("emoji") == emoji - if emoji and not toggling_off: + if emoji and not (previous is not None and previous.get("emoji") == emoji): reactions.append({"emoji": _scrub_surrogates(emoji), "author": author, "at": time.time()}) if reactions: meta[self.REACTIONS_METADATA_KEY] = reactions @@ -500,8 +462,7 @@ class SessionMessagesMixin: meta.pop(self.REACTIONS_METADATA_KEY, None) conn.execute( "UPDATE messages SET display_metadata = ? WHERE id = ?", - (self._encode_display_metadata(meta) if meta else None, message_row_id), - ) + (self._encode_display_metadata(meta) if meta else None, message_row_id)) return reactions return self._execute_write(_do) @@ -512,24 +473,21 @@ class SessionMessagesMixin: return [] row = self._read_one( "SELECT display_metadata FROM messages WHERE id = ? AND session_id = ?", - (message_row_id, session_id), - ) + (message_row_id, session_id)) return self._reaction_list(self._decode_display_metadata(row[0])) if row is not None else [] def take_unseen_reactions(self, session_id: str, *, author: str = "user") -> List[Dict[str, Any]]: """Return *author*'s not-yet-surfaced reactions and mark them seen. Reactions are announced on the NEXT user turn (never by rewriting the reacted - message — cache-safe); the ``seen`` stamp makes each announcement exactly once. - """ + message — cache-safe); the ``seen`` stamp makes each announcement exactly once.""" if not session_id: return [] def _do(conn): rows = conn.execute( "SELECT id, role, content, display_metadata FROM messages " - "WHERE session_id = ? AND active = 1 AND display_metadata IS NOT NULL " - "ORDER BY id", + "WHERE session_id = ? AND active = 1 AND display_metadata IS NOT NULL ORDER BY id", (session_id,), ).fetchall() pending = [] @@ -548,16 +506,12 @@ class SessionMessagesMixin: changed = True content = self._decode_content(row["content"]) pending.append({ - "row_id": row["id"], - "role": row["role"], - "emoji": reaction.get("emoji") or "", - "text": content if isinstance(content, str) else "", - }) + "row_id": row["id"], "role": row["role"], "emoji": reaction.get("emoji") or "", + "text": content if isinstance(content, str) else ""}) if changed: conn.execute( "UPDATE messages SET display_metadata = ? WHERE id = ?", - (self._encode_display_metadata(meta), row["id"]), - ) + (self._encode_display_metadata(meta), row["id"])) return pending return self._execute_write(_do) or [] @@ -565,20 +519,16 @@ class SessionMessagesMixin: def latest_message_row_id( self, session_id: str, *, role: str = "user", offset: int = 0, require_text: bool = True ) -> Optional[int]: - """Row id of the most recent active message with *role*, or ``None``. - - ``offset`` steps to earlier turns (1 = the one before the latest). ``require_text`` - skips rows without plain-text content (tool-call-only turns, attachment stubs) - so "the latest message" never resolves to an invisible bubble. - """ + """Row id of the most recent active *role* message, or ``None``. ``offset`` steps + to earlier turns; ``require_text`` skips rows without plain-text content so "the + latest message" never resolves to an invisible bubble.""" if not session_id or role not in {"user", "assistant"} or offset < 0: return None text_filter = "AND content IS NOT NULL AND TRIM(content) != '' " if require_text else "" row = self._read_one( "SELECT id FROM messages WHERE session_id = ? AND role = ? " f"AND active = 1 {text_filter}ORDER BY id DESC LIMIT 1 OFFSET ?", - (session_id, role, int(offset)), - ) + (session_id, role, int(offset))) return row[0] if row else None def latest_user_message_row_id(self, session_id: str) -> Optional[int]: @@ -592,17 +542,13 @@ class SessionMessagesMixin: return None row = self._read_one( "SELECT role FROM messages WHERE id = ? AND session_id = ? AND active = 1", - (int(row_id), session_id), - ) + (int(row_id), session_id)) return row[0] if row else None def _insert_message_rows(self, conn, session_id: str, messages: List[Dict[str, Any]]) -> tuple[int, int]: - """Insert *messages* as fresh active rows inside the caller's write txn. - - Returns ``(inserted_count, tool_call_count)``; does NOT touch sessions.* - counters (callers reconcile them differently). Reasoning columns are kept for - assistant rows only. Stamps ``msg["_row_id"]``. - """ + """Insert *messages* as fresh active rows inside the caller's write txn. Returns + ``(inserted, tool_call_count)``; never touches sessions.* counters (callers + reconcile differently). Reasoning columns kept for assistant rows only.""" now_ts = time.time() inserted = tool_calls_total = 0 for msg in messages: @@ -613,8 +559,7 @@ class SessionMessagesMixin: _INSERT_MESSAGE_SQL, self._message_row_params( session_id, role, msg, tool_calls, message_timestamp, keep_reasoning=role == "assistant", - ), - ) + )) if isinstance(msg, dict) and cur.lastrowid is not None: msg["_row_id"] = cur.lastrowid inserted += 1 @@ -626,18 +571,12 @@ class SessionMessagesMixin: self, session_id: str, messages: List[Dict[str, Any]], active_only: bool = False, archive_dropped: bool = False, reject_active_turn_lease: bool = False, ) -> None: - """Atomically replace the stored messages for a session (/retry, /undo, /compress). - - DESTRUCTIVE by default: every row is DELETEd (and leaves the FTS index). - ``active_only=True`` replaces only ``active = 1`` rows, leaving soft-archived - rows (compacted turns, rewind rows) untouched — required when sharing a session - id with an agent doing in-place compaction. ``archive_dropped=True`` SOFT-archives - the live rows (``active = 0, compacted = 0``, rewind-style) instead of deleting: - the mode rewind/edit/regenerate must use, since a DELETE leaves nothing to - recover from; it implies active-only handling. ``reject_active_turn_lease=True`` - runs the lease check in the same write txn for user-initiated rewrites that do - not own the cross-process lease. - """ + """Atomically replace a session's stored messages (/retry, /undo, /compress). + DESTRUCTIVE by default (rows DELETEd, leave FTS). ``active_only`` spares + soft-archived rows (needed alongside in-place compaction). ``archive_dropped`` + SOFT-archives the live rows rewind-style instead of deleting — what rewind/edit/ + regenerate must use, since a DELETE leaves nothing to recover. ``reject_active_ + turn_lease`` runs the lease check in-txn for user rewrites that don't own it.""" from hermes_state import CompressionSessionClosedError active_clause = " AND active = 1" if active_only else "" @@ -662,8 +601,7 @@ class SessionMessagesMixin: total_messages, total_tool_calls = self._insert_message_rows(conn, session_id, messages) conn.execute( "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", - (total_messages, total_tool_calls, session_id), - ) + (total_messages, total_tool_calls, session_id)) self._execute_write(_do) @@ -680,9 +618,7 @@ class SessionMessagesMixin: 0 for an empty/unknown session.""" if not session_id: return 0 - row = self._read_one( - "SELECT COALESCE(MAX(id), 0) FROM messages WHERE session_id = ? AND active = 1", (session_id,), - ) + row = self._read_one("SELECT COALESCE(MAX(id), 0) FROM messages WHERE session_id = ? AND active = 1", (session_id,)) return int(row[0]) if row else 0 def _tail_rows_after_watermark(self, conn, sql: str, params) -> Tuple[List[int], int]: @@ -695,61 +631,40 @@ class SessionMessagesMixin: """Pure-SQL column clone of *tail_ids* as fresh live rows (new id, active=1, compacted=0, everything else byte-exact; FTS triggers index the clones). With *session_id* the clones land in that session instead of the originals'.""" - skip = ("id", "active", "compacted") + (("session_id",) if session_id is not None else ()) + retarget = session_id is not None + skip = ("id", "active", "compacted") + (("session_id",) if retarget else ()) col_list = ", ".join(c for c in self._message_column_names(conn) if c not in skip) - placeholders = _placeholders(tail_ids) - if session_id is None: - conn.execute( - f"INSERT INTO messages ({col_list}, active, compacted) " - f"SELECT {col_list}, 1, 0 FROM messages " - f"WHERE id IN ({placeholders}) ORDER BY id", - tail_ids, - ) - else: - conn.execute( - f"INSERT INTO messages ({col_list}, session_id, active, compacted) " - f"SELECT {col_list}, ?, 1, 0 FROM messages " - f"WHERE id IN ({placeholders}) ORDER BY id", - [session_id, *tail_ids], - ) + conn.execute( + f"INSERT INTO messages ({col_list}, {'session_id, ' if retarget else ''}active, compacted) " + f"SELECT {col_list}, {'?, ' if retarget else ''}1, 0 FROM messages " + f"WHERE id IN ({_placeholders(tail_ids)}) ORDER BY id", + [session_id, *tail_ids] if retarget else tail_ids) def archive_and_compact( self, session_id: str, compacted_messages: List[Dict[str, Any]], model_config_patch: Optional[Dict[str, Any]] = None, watermark: Optional[int] = None, lock_holder: Optional[str] = None, tail_count: int = 0, ) -> int: - """Non-destructive in-place compaction under ONE durable session id. + """Non-destructive in-place compaction under ONE durable session id: soft-archive + the active rows (``active=0, compacted=1`` — "summarized away", still searchable) + and insert *compacted_messages* as fresh active rows, atomically. - Soft-archives the active rows (``active=0, compacted=1`` — "summarized away", - still found by search_messages and readable with include_inactive) and inserts - *compacted_messages* as fresh active rows, atomically. Live-context loads - filter ``active = 1`` so the model reloads only the compacted set. - - *watermark* (``get_active_message_watermark`` at compression START): rows with - ``id > watermark`` arrived during the slow summary and are re-sequenced after - the compacted set by a pure-SQL column clone (fresh ids — consumers re-resolve - by content); ``None`` archives everything. *lock_holder*: the commit verifies - inside the txn that the compression lock is still held and unexpired, so a - reclaimed lease fails instead of clobbering the winner. *tail_count*: the LAST - N rows of *compacted_messages* are the verbatim carried-forward tail; their - originals (at/below the watermark) and the watermark clones' originals are - superseded duplicates and get rewind-style flags (``active=0, compacted=0``) so - session_search doesn't return each carried message once per compaction. - - ``message_count`` becomes the ACTIVE count; ``model_config_patch`` merges into - the session JSON in the same txn (``None`` value removes a key). Returns the new - active count. - """ + *watermark* (captured at compression START): rows with ``id > watermark`` arrived + during the slow summary and are re-sequenced after the compacted set by a pure- + SQL column clone (fresh ids); ``None`` archives everything. *lock_holder*: the + commit verifies in-txn that the lease is still held, so a reclaimed lease fails + instead of clobbering the winner. *tail_count*: the LAST N compacted rows are the + verbatim carried-forward tail; their originals and the watermark clones' originals + are superseded duplicates and get rewind-style flags (``active=0, compacted=0``) + so search doesn't return each carried message once per compaction. + ``message_count`` becomes the ACTIVE count; ``model_config_patch`` merges in the + same txn (``None`` removes a key). Returns the new active count.""" from hermes_state import SessionCompressionInProgressError def _do(conn): if lock_holder is not None: lock_row = conn.execute(_COMPRESSION_LOCK_ROW_SQL, (session_id,)).fetchone() - if ( - lock_row is None - or lock_row["holder"] != lock_holder - or float(lock_row["expires_at"]) <= time.time() - ): + if lock_row is None or lock_row["holder"] != lock_holder or float(lock_row["expires_at"]) <= time.time(): raise SessionCompressionInProgressError( f"Compression lease for {session_id!r} lost before " "commit; refusing to publish a stale compaction" @@ -765,22 +680,17 @@ class SessionMessagesMixin: tail_tool_calls = 0 if watermark is not None: tail_ids, tail_tool_calls = self._tail_rows_after_watermark( - conn, - "SELECT id, tool_calls FROM messages " - "WHERE session_id = ? AND active = 1 AND id > ? " - "ORDER BY id", - (session_id, int(watermark)), - ) + conn, "SELECT id, tool_calls FROM messages " + "WHERE session_id = ? AND active = 1 AND id > ? ORDER BY id", + (session_id, int(watermark))) # Rewind targets sit AT/BELOW the watermark (the compressor only saw rows up # to it); without the bound a concurrent append would steal a LIMIT slot. - rewind_tail_ids: Optional[list[int]] = None + rewind_ids: list[int] = [] if tail_count > 0: if watermark is not None: tail_rows = conn.execute( - "SELECT id FROM messages " - "WHERE session_id = ? AND active = 1 AND id <= ? " - "ORDER BY id DESC LIMIT ?", - (session_id, int(watermark), int(tail_count)), + "SELECT id FROM messages WHERE session_id = ? AND active = 1 AND id <= ? " + "ORDER BY id DESC LIMIT ?", (session_id, int(watermark), int(tail_count)), ).fetchall() else: tail_rows = conn.execute( @@ -788,27 +698,21 @@ class SessionMessagesMixin: "WHERE session_id = ? AND active = 1 ORDER BY id DESC LIMIT ?", (session_id, int(tail_count)), ).fetchall() - rewind_tail_ids = [int(row["id"]) for row in tail_rows] - rewind_ids = [*(rewind_tail_ids or []), *tail_ids] + rewind_ids = [int(row["id"]) for row in tail_rows] + rewind_ids += tail_ids if rewind_ids: placeholders = _placeholders(rewind_ids) conn.execute( "UPDATE messages SET active = 0, compacted = 0 " - f"WHERE session_id = ? AND id IN ({placeholders})", - [session_id, *rewind_ids], - ) + f"WHERE session_id = ? AND id IN ({placeholders})", [session_id, *rewind_ids]) conn.execute( "UPDATE messages SET active = 0, compacted = 1 " "WHERE session_id = ? AND active = 1 " - f"AND id NOT IN ({placeholders})", - [session_id, *rewind_ids], - ) + f"AND id NOT IN ({placeholders})", [session_id, *rewind_ids]) else: conn.execute( "UPDATE messages SET active = 0, compacted = 1 " - "WHERE session_id = ? AND active = 1", - (session_id,), - ) + "WHERE session_id = ? AND active = 1", (session_id,)) inserted, tool_calls_total = self._insert_message_rows(conn, session_id, compacted_messages) if tail_ids: self._clone_message_rows(conn, tail_ids) @@ -817,14 +721,12 @@ class SessionMessagesMixin: if model_config_patch is None: conn.execute( "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", - (inserted, tool_calls_total, session_id), - ) + (inserted, tool_calls_total, session_id)) else: conn.execute( "UPDATE sessions SET message_count = ?, tool_call_count = ?, " "model_config = ? WHERE id = ?", - (inserted, tool_calls_total, patched_model_config, session_id), - ) + (inserted, tool_calls_total, patched_model_config, session_id)) return inserted return self._execute_write(_do) @@ -839,53 +741,39 @@ class SessionMessagesMixin: return cols def set_latest_user_api_content(self, session_id: str, content: Any, api_content: str) -> int: - """Backfill the ``api_content`` sidecar onto the newest ACTIVE user row. - - In-place preflight compaction inserts the current user row BEFORE the turn - prologue composes the sidecar, and the later persist identity-skips compacted - dicts; without this the reload would replay clean content and reopen the - prompt-cache divergence. The ``content`` match guards against a racing rewrite. - Returns rows updated (0 or 1). - """ + """Backfill the ``api_content`` sidecar onto the newest ACTIVE user row. Preflight + compaction inserts that row BEFORE the sidecar is composed and the later persist + identity-skips compacted dicts; without this a reload would reopen the prompt- + cache divergence. The ``content`` match guards a racing rewrite. Returns 0/1.""" from hermes_state import _scrub_surrogates return self._write_rowcount( - "UPDATE messages SET api_content = ? WHERE id = (" - "SELECT id FROM messages " - "WHERE session_id = ? AND role = 'user' AND active = 1 " - "ORDER BY id DESC LIMIT 1" + "UPDATE messages SET api_content = ? WHERE id = (SELECT id FROM messages " + "WHERE session_id = ? AND role = 'user' AND active = 1 ORDER BY id DESC LIMIT 1" ") AND content IS ?", - (_scrub_surrogates(api_content), session_id, self._encode_content(content)), - ) + (_scrub_surrogates(api_content), session_id, self._encode_content(content))) def _dedupe_display_generations(self, rows): - """Collapse compaction generations so each logical message appears once. - - Compaction copies the protected tail into each generation (same - role/content/timestamp, different ``active``/id); prefer the live row, then the - newest generation. The ONE definition shared by every display projection - (get_messages, get_resume_conversations, get_ancestor_display_prefix, - get_messages_as_conversation). *rows* must be ordered by ``id``; order is kept. - """ + """Collapse compaction generations so each logical message appears once: the + protected tail is copied into each generation (same role/content/timestamp, + different ``active``/id); prefer the live row, then the newest. The ONE definition + shared by every display projection. *rows* must be ordered by ``id``.""" seen: Dict[Tuple[Any, ...], Any] = {} for row in rows: dedupe_content = row["content"] if row["role"] == "user": from agent.context_compressor import split_user_originated_turn - handoff, live_view = split_user_originated_turn({ "role": "user", "content": self._decode_content(row["content"]), "display_kind": row["display_kind"], - "display_metadata": self._decode_display_metadata(row["display_metadata"]), - }) + "display_metadata": self._decode_display_metadata(row["display_metadata"])}) if handoff is not None and live_view is not None: dedupe_content = self._encode_content(live_view.get("content")) # Tool fields are part of the key: identical tool messages across generations # collapse, distinct tool calls sharing role/content/timestamp never merge. key = ( row["role"], dedupe_content, row["timestamp"], - row["tool_call_id"], row["tool_calls"], row["tool_name"], - ) + row["tool_call_id"], row["tool_calls"], row["tool_name"]) cur = seen.get(key) if cur is None or (row["active"], row["id"]) > (cur["active"], cur["id"]): seen[key] = row @@ -901,8 +789,7 @@ class SessionMessagesMixin: msg["content"] = self._decode_content(msg["content"]) if msg.get("tool_calls"): msg["tool_calls"] = _json_or( - msg["tool_calls"], [], - f"Failed to deserialize tool_calls in {warn_context}, falling back to []", + msg["tool_calls"], [], f"Failed to deserialize tool_calls in {warn_context}, falling back to []", ) if msg.get("display_metadata") is not None: msg["display_metadata"] = self._decode_display_metadata(msg["display_metadata"]) @@ -921,15 +808,10 @@ class SessionMessagesMixin: limit: Optional[int] = None, offset: int = 0, latest: bool = False, after_id: Optional[int] = None, ) -> List[Dict[str, Any]]: - """Load messages for a session in insertion order (AUTOINCREMENT id, never - timestamp — clocks regress on WSL2/NTP steps). - - ``include_inactive`` loads soft-deleted rewind rows; ``include_compacted`` adds - rows preserved by in-place compaction (durable display history a transcript - read must not drop) but not rewind rows. ``limit``/``offset`` page; ``latest`` - measures the offset back from the newest row and still returns chronological - order. ``after_id`` is keyset paging (``id > after_id``), ascending only. - """ + """Load a session's messages in insertion order (id, never timestamp — clocks + regress). ``include_inactive``: rewind rows too; ``include_compacted``: compaction- + archived display history (not rewind rows). ``latest`` pages back from the newest + row but still returns chronological order; ``after_id`` is keyset paging.""" if after_id is not None and (latest or offset): raise ValueError("after_id is incompatible with latest/offset paging") if after_id is not None and include_compacted: @@ -940,8 +822,7 @@ class SessionMessagesMixin: # dedupe generations, then page. rows = self._dedupe_display_generations(self._read_all( "SELECT * FROM messages WHERE session_id = ?" + active_clause + " ORDER BY id ASC", - [session_id], - )) + [session_id])) if latest: rows = rows[::-1] rows = rows[offset:] @@ -955,9 +836,7 @@ class SessionMessagesMixin: "SELECT * FROM messages WHERE session_id = ?" f"{active_clause}{keyset_clause} ORDER BY id {'DESC' if latest else 'ASC'}" ) - params: list = [session_id] - if after_id is not None: - params.append(after_id) + params: list = [session_id] if after_id is None else [session_id, after_id] if limit is not None or offset: # SQLite's OFFSET requires LIMIT; -1 means "no limit". sql += " LIMIT ? OFFSET ?" @@ -974,64 +853,44 @@ class SessionMessagesMixin: ids = [s for s in session_ids if s] for start in range(0, len(ids), 900): # SQLite's bound-variable ceiling. chunk = ids[start : start + 900] - rows = self._read_all( + found.extend({"session_id": row[0], "content": row[1]} for row in self._read_all( f"""SELECT session_id, content FROM messages - WHERE session_id IN ({",".join("?" * len(chunk))}) + WHERE session_id IN ({_placeholders(chunk)}) AND role = 'tool' AND content LIKE '%/pull/%' ORDER BY id ASC""", chunk, - ) - found.extend({"session_id": row[0], "content": row[1]} for row in rows) + )) return found def get_messages_around(self, session_id: str, around_message_id: int, window: int = 5) -> Dict[str, Any]: - """Window of up to *window* messages either side of an anchor id (id ascending). - - ``messages_before``/``messages_after`` count rows in the returned slice strictly - before/after the anchor; less than *window* means a session boundary. Empty - window when the anchor is not a row of *session_id*. - """ + """Up to *window* messages either side of an anchor id (ascending). ``messages_ + before``/``_after`` count the slice strictly around the anchor (fewer than + *window* = session boundary). Empty when the anchor is not in *session_id*.""" window = max(window, 0) with self._read_ctx() as conn: - anchor_exists = conn.execute( - "SELECT 1 FROM messages WHERE id = ? AND session_id = ? LIMIT 1", - (around_message_id, session_id), - ).fetchone() - if not anchor_exists: + if not conn.execute( + "SELECT 1 FROM messages WHERE id = ? AND session_id = ? LIMIT 1", (around_message_id, session_id), + ).fetchone(): return {"window": [], "messages_before": 0, "messages_after": 0} before_rows = conn.execute( - "SELECT * FROM messages " - "WHERE session_id = ? AND id <= ? " - "ORDER BY id DESC LIMIT ?", + "SELECT * FROM messages WHERE session_id = ? AND id <= ? ORDER BY id DESC LIMIT ?", (session_id, around_message_id, window + 1), ).fetchall() after_rows = conn.execute( - "SELECT * FROM messages " - "WHERE session_id = ? AND id > ? " - "ORDER BY id ASC LIMIT ?", + "SELECT * FROM messages WHERE session_id = ? AND id > ? ORDER BY id ASC LIMIT ?", (session_id, around_message_id, window), ).fetchall() rows = list(reversed(before_rows)) + list(after_rows) - return { - "window": [ - self._row_to_message_dict(row, warn_context="get_messages_around", summary_flag=False) - for row in rows - ], - "messages_before": max(0, len(before_rows) - 1), # before_rows includes the anchor - "messages_after": len(after_rows), - } + window_msgs = [self._row_to_message_dict(r, warn_context="get_messages_around", summary_flag=False) for r in rows] + # before_rows includes the anchor itself. + return {"window": window_msgs, "messages_before": max(0, len(before_rows) - 1), "messages_after": len(after_rows)} def resolve_resume_session_id(self, session_id: str) -> str: - """Redirect a resume target to the descendant session that holds the messages. - - Follows the compression chain to the live tip first (``get_compression_tip`` is - lineage-aware: only children of compression-ended parents, so delegation/branch - children never hijack the resume), then walks ``parent_session_id`` forward, - returning the deepest node with messages — never short-circuiting on the start - node, since a continuation may hold the newer turns. Branch, delegate, reset - and tool children are skipped (they carry ``parent_session_id`` too). Returns - *session_id* unchanged when nothing has messages. Depth cap 32. - """ + """Redirect a resume target to the descendant that holds the messages: follow the + compression chain to the live tip (lineage-aware, so delegate/branch children never + hijack it), then walk ``parent_session_id`` forward to the DEEPEST node with + messages (a continuation may hold newer turns), skipping branch/delegate/reset/ + tool children. Unchanged when nothing has messages. Depth cap 32.""" if not session_id: return session_id try: @@ -1046,24 +905,18 @@ class SessionMessagesMixin: best = None # deepest node with messages for _ in range(32): try: - row = conn.execute( + if conn.execute( "SELECT 1 FROM messages WHERE session_id = ? LIMIT 1", (current,), - ).fetchone() - except Exception: - return session_id - if row is not None: - best = current - try: + ).fetchone() is not None: + best = current child_row = conn.execute( - "SELECT id FROM sessions AS child " - "WHERE child.parent_session_id = ? " + "SELECT id FROM sessions AS child WHERE child.parent_session_id = ? " " AND json_extract(COALESCE(child.model_config, '{}'), '$._branched_from') IS NULL " " AND json_extract(COALESCE(child.model_config, '{}'), '$._delegate_from') IS NULL " " AND json_extract(COALESCE(child.model_config, '{}'), '$._reset_from') IS NULL " f" AND NOT {_legacy_reset_child_sql('child', _RESET_END_REASONS_SQL)} " " AND COALESCE(child.source, '') != 'tool' " - "ORDER BY child.started_at DESC, child.id DESC LIMIT 1", - (current,), + "ORDER BY child.started_at DESC, child.id DESC LIMIT 1", (current,), ).fetchone() except Exception: return session_id @@ -1084,8 +937,7 @@ class SessionMessagesMixin: return conn.execute( f"{prefix}{self._CONVERSATION_ROW_COLUMNS} " f"FROM messages WHERE session_id IN ({_placeholders(session_ids)})" - f"{active_clause} ORDER BY id", - tuple(session_ids), + f"{active_clause} ORDER BY id", tuple(session_ids), ).fetchall() def get_messages_as_conversation( @@ -1093,15 +945,11 @@ class SessionMessagesMixin: repair_alternation: bool = False, include_row_ids: bool = False, include_compacted: bool = False, ) -> List[Dict[str, Any]]: - """Load messages in OpenAI conversation format (gateway history restore). - - ``include_compacted`` adds compaction-archived rows deduped by - :meth:`_dedupe_display_generations` — DISPLAY reads only; the model-fed restore - must not pass it or a resume regrows the history compaction summarized away. - ``repair_alternation`` runs ``repair_message_sequence`` on the loaded list - (LIVE REPLAY callers) so a durable ``user;user`` pair doesn't re-trigger the - per-request repair forever; the stored transcript is never mutated. - """ + """Load messages in OpenAI conversation format. ``include_compacted`` (deduped + display history) is for DISPLAY reads only — the model-fed restore must not regrow + what compaction summarized away. ``repair_alternation`` repairs the loaded list for + LIVE REPLAY callers (a durable ``user;user`` pair would otherwise re-trigger the + per-request repair forever); the stored transcript is never mutated.""" session_ids = [session_id] if include_ancestors and not self._is_explicit_branch_session(session_id): session_ids = self._session_lineage_root_to_tip(session_id) @@ -1112,18 +960,13 @@ class SessionMessagesMixin: rows = self._dedupe_display_generations(rows) return self._rows_to_conversation( rows, session_id=session_id, include_ancestors=include_ancestors, - repair_alternation=repair_alternation, include_row_ids=include_row_ids, - ) + repair_alternation=repair_alternation, include_row_ids=include_row_ids) def _dedupe_replayed_user(self, messages, msg, exact_user_clones) -> Tuple[bool, Any]: - """Ancestor-lineage dedupe for one decoded user *msg*. - - Returns ``(skip, exact_clone_key)``. Watermark rotation column-clones the - concurrent tail into the child after the summary, so the copies need not be - adjacent: an exact ``(timestamp, canonical content)`` clone index is checked - first, then the adjacent-duplicate heuristic. A rotated child carrier wins over - the simpler ancestor copy (it owns the durable row id and the summary scaffold). - """ + """Ancestor-lineage dedupe for one decoded user *msg* -> ``(skip, exact_clone_key)``. + Rotation column-clones the concurrent tail into the child, so copies need not be + adjacent: the exact ``(timestamp, canonical content)`` clone index is checked first, + then the adjacent heuristic. A rotated child carrier wins over the ancestor copy.""" canonical_content, _is_composite = self._canonical_replayed_user_content(msg) exact_clone_key = self._exact_replayed_user_clone_key(msg.get("timestamp"), canonical_content) previous_exact = exact_user_clones.get(exact_clone_key) if exact_clone_key is not None else None @@ -1151,10 +994,8 @@ class SessionMessagesMixin: Every dict is stamped ``_DB_PERSISTED_MARKER_KEY`` at the source (born durable) so an identity-losing handoff never re-appends the whole transcript on flush. - ``_row_id`` is opt-in (gateway reactions). ``api_content`` is returned VERBATIM - (no sanitize/strip): the replay path substitutes it to keep the provider prompt - cache byte-stable. Reasoning fields are restored on assistant rows only. - """ + ``_row_id`` is opt-in (gateway reactions). Reasoning fields are restored on + assistant rows only. Key order of each dict is stable (see the column tables).""" from hermes_state import _strip_background_review_harness, _strip_stale_tool_call_markers messages = [] exact_user_clones: Dict[Tuple[Any, str], Dict[str, Any]] = {} @@ -1162,46 +1003,37 @@ class SessionMessagesMixin: content = self._decode_content(row["content"]) if row["role"] in {"user", "assistant"} and isinstance(content, str): content = sanitize_context(content).strip() - msg = {"role": row["role"], "content": content} - # Underscore-prefixed like ``_row_id``: every transport strips it before the - # wire, and compression's assembly copies deliberately strip it so rotated - # child handoffs still flush (see _fresh_compaction_message_copy). - msg[_DB_PERSISTED_MARKER_KEY] = True + # The persisted marker is underscore-prefixed like ``_row_id``: every transport + # strips it before the wire, and compression's assembly copies deliberately + # strip it so rotated child handoffs still flush (see _fresh_compaction_message_copy). + msg = {"role": row["role"], "content": content, _DB_PERSISTED_MARKER_KEY: True} if include_row_ids and row["id"] is not None: msg["_row_id"] = row["id"] - for col in ("api_content", "display_kind"): - if row[col]: - msg[col] = row[col] + msg.update((col, row[col]) for col in _VERBATIM_COLS if row[col]) if row["display_metadata"]: decoded = self._decode_display_metadata(row["display_metadata"]) if decoded is not None: msg["display_metadata"] = decoded if include_summary_markers and row["_compressed_summary"]: msg["_compressed_summary"] = True - for col in ("timestamp", "tool_call_id", "tool_name", "effect_disposition"): - if row[col]: - msg[col] = row[col] + msg.update((col, row[col]) for col in _META_COLS if row[col]) if row["tool_calls"]: msg["tool_calls"] = _json_or( row["tool_calls"], [], - "Failed to deserialize tool_calls in conversation replay, falling back to []", - ) + "Failed to deserialize tool_calls in conversation replay, falling back to []") # Platform-side id exposed as ``message_id`` (JSONL transcript compat). if row["platform_message_id"]: msg["message_id"] = row["platform_message_id"] if row["observed"]: msg["observed"] = True if row["role"] == "assistant": - for col in ("finish_reason", "reasoning"): - if row[col]: - msg[col] = row[col] + msg.update((col, row[col]) for col in ("finish_reason", "reasoning") if row[col]) if row["reasoning_content"] is not None: msg["reasoning_content"] = row["reasoning_content"] - for col in ("reasoning_details", "codex_reasoning_items", "codex_message_items"): - if row[col]: - msg[col] = _json_or( - row[col], None, f"Failed to deserialize {col}, falling back to None", - ) + msg.update( + (col, _json_or(row[col], None, f"Failed to deserialize {col}, falling back to None")) + for col in _ASSISTANT_JSON_COLS if row[col] + ) exact_clone_key = None if include_ancestors: skip, exact_clone_key = self._dedupe_replayed_user(messages, msg, exact_user_clones) @@ -1217,48 +1049,36 @@ class SessionMessagesMixin: messages = _strip_stale_tool_call_markers(messages) if repair_alternation and messages: from agent.agent_runtime_helpers import repair_message_sequence - repaired = repair_message_sequence(None, messages) if repaired: logger.info( "Repaired %d message-alternation violation(s) while " "restoring session %s — durable transcript kept them, " - "see repair_message_sequence", - repaired, - session_id, - ) + "see repair_message_sequence", repaired, session_id) return messages def get_resume_conversations(self, session_id: str) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]: - """``(model_history, display_history)`` for a session resume from ONE SELECT. - - ``model_history``: the tip's active rows, alternation-repaired, with the summary - marker kept for pre-compress checkpointing. ``display_history``: the full - compression lineage (``/branch`` sessions are their own lineage) verbatim, with - compaction-archived rows included and deduped, plus replayed-user dedup. Byte- - identical to the separate reads (test_get_resume_conversations_matches_separate_reads). - """ + """``(model_history, display_history)`` for a resume from ONE SELECT. model: the + tip's active rows, alternation-repaired, summary marker kept for pre-compress + checkpointing. display: the full lineage (``/branch`` sessions stand alone) with + compaction-archived rows deduped. Byte-identical to the separate reads.""" session_ids = self._resume_lineage_ids(session_id) rows = self._fetch_conversation_rows(session_ids, _DISPLAY_ACTIVE_CLAUSE, with_session_id=True) # The model projection stays active-only: it is the compressed working context. tip_rows = [r for r in rows if r["session_id"] == session_id and r["active"]] model_history = self._rows_to_conversation( tip_rows, session_id=session_id, include_ancestors=False, repair_alternation=True, - include_row_ids=True, include_summary_markers=True, - ) + include_row_ids=True, include_summary_markers=True) display_history = self._rows_to_conversation( self._dedupe_display_generations(rows), session_id=session_id, - include_ancestors=True, repair_alternation=False, include_row_ids=True, - ) + include_ancestors=True, repair_alternation=False, include_row_ids=True) return model_history, display_history def _resume_lineage_ids(self, session_id: str) -> List[str]: """Session ids a full (display) resume materializes: the compression lineage, or the session alone for an explicit ``/branch`` copy. Shared by the resume readers and the resume guard so the guard counts exactly what a resume loads.""" - if self._is_explicit_branch_session(session_id): - return [session_id] - return self._session_lineage_root_to_tip(session_id) + return [session_id] if self._is_explicit_branch_session(session_id) else self._session_lineage_root_to_tip(session_id) def _resume_count_scope(self, session_id: str, tip_only: bool) -> Tuple[List[str], str]: """``tip_only``: the tip's ACTIVE rows (model restore); else the full-lineage @@ -1271,21 +1091,15 @@ class SessionMessagesMixin: """Count the rows a resume would materialize (see ``_resume_count_scope``).""" session_ids, active_clause = self._resume_count_scope(session_id, tip_only) row = self._read_one( - f"SELECT COUNT(*) FROM messages " - f"WHERE session_id IN ({_placeholders(session_ids)}) AND {active_clause}", - tuple(session_ids), - ) + f"SELECT COUNT(*) FROM messages WHERE session_id IN ({_placeholders(session_ids)}) AND {active_clause}", + tuple(session_ids)) return int(row[0] if row else 0) def assert_resume_safe(self, session_id: str, max_messages: Optional[int] = None, *, tip_only: bool = False) -> int: - """Return the resume row count or raise ``SessionResumeTooLargeError``. - - ``max_messages=None`` reads ``sessions.max_resume_messages``; 0 disables the - guard and returns 0 without counting. ``tip_only`` bounds only the tip's active - rows, for callers that never materialize the lineage in memory — a heavily - compressed conversation (~29k lineage rows behind a ~700-row tip) is exactly - what compression should produce and must not be rejected. - """ + """Resume row count, or raise ``SessionResumeTooLargeError``. ``max_messages=None`` + reads config; 0 disables the guard without counting. ``tip_only`` bounds only the + tip's active rows for callers that never materialize the lineage — a heavily + compressed conversation is what compression should produce, not a rejection.""" from hermes_state import SessionResumeTooLargeError, resolved_max_resume_messages if max_messages is None: max_messages = resolved_max_resume_messages() @@ -1298,24 +1112,18 @@ class SessionMessagesMixin: "SELECT COUNT(*) FROM (" f"SELECT 1 FROM messages WHERE session_id IN ({_placeholders(session_ids)}) " f"AND {active_clause} LIMIT ?" - ")", - (*session_ids, max_messages + 1), - ) + ")", (*session_ids, max_messages + 1)) message_count = int(row[0] if row else 0) if message_count > max_messages: raise SessionResumeTooLargeError( message_count, max_messages, - scope="in its tip segment" if tip_only else "across its lineage", - ) + scope="in its tip segment" if tip_only else "across its lineage") return message_count def get_ancestor_display_prefix(self, session_id: str) -> List[Dict[str, Any]]: - """Ancestor-only display messages of a lineage (rows with ``session_id !=`` tip). - - ``session.resume`` prepends this to the live model history. Identifying - ancestors by row origin (not ``display[:len(display) - len(model)]``) avoids - overcounting when alternation repair removes tip messages from the middle. - """ + """Ancestor-only display messages of a lineage (row ``session_id != tip``), which + ``session.resume`` prepends to the model history. Identified by row origin, not + ``display[:len(display) - len(model)]``, so alternation repair cannot overcount.""" session_ids = self._resume_lineage_ids(session_id) if len(session_ids) <= 1: return [] @@ -1328,13 +1136,10 @@ class SessionMessagesMixin: lineage = self._rows_to_conversation( rows, session_id=session_id, include_ancestors=True, repair_alternation=False, include_row_ids=True, ) - prefix: List[Dict[str, Any]] = [] - for message in lineage: - if message.get("_row_id") in ancestor_ids: - projected = message.copy() - projected.pop("_row_id", None) - prefix.append(projected) - return prefix + return [ + {k: v for k, v in message.items() if k != "_row_id"} + for message in lineage if message.get("_row_id") in ancestor_ids + ] def get_conversation_root(self, session_id: str) -> str: """ROOT id of *session_id*'s lineage — the stable conversation id across @@ -1349,13 +1154,11 @@ class SessionMessagesMixin: if msg.get("role") != "user": return None, False from agent.context_compressor import split_user_originated_turn - handoff, live_view = split_user_originated_turn(msg) is_composite = handoff is not None and live_view is not None return ( live_view.get("content") if is_composite and live_view is not None else msg.get("content"), - is_composite, - ) + is_composite) @staticmethod def _exact_replayed_user_clone_key(timestamp: Any, content: Any) -> Optional[Tuple[Any, str]]: @@ -1372,13 +1175,10 @@ class SessionMessagesMixin: def _find_duplicate_replayed_user_message( messages: List[Dict[str, Any]], msg: Dict[str, Any] ) -> Optional[Tuple[int, bool]]: - """Return an adjacent replay duplicate and whether *msg* must win. - - Rotation may persist the current ask once in the parent and again inside a - composite child carrier; compare the canonical live payload for carriers while - keeping the exact-string dedupe for ordinary replayed users. The child carrier - wins (it owns the durable row id and the retained scaffold). - """ + """Adjacent replay duplicate ``(index, prefer_current)`` or None. Rotation may + persist the current ask in the parent and again inside a composite child carrier: + carriers compare by canonical live payload, ordinary users by exact string. The + child carrier wins (it owns the durable row id and the retained scaffold).""" from hermes_state import SessionDB if msg.get("role") != "user": return None @@ -1409,32 +1209,15 @@ class SessionMessagesMixin: rows = conn.execute( "SELECT tool_calls FROM messages WHERE session_id = ? AND active = 1", (session_id,), ).fetchall() - tool_call_count = 0 - for row in rows: - raw = row[0] - if not raw: - continue - try: - decoded = json.loads(raw) if isinstance(raw, str) else raw - except (json.JSONDecodeError, TypeError): - continue - if isinstance(decoded, list): - tool_call_count += len(decoded) - elif decoded: - tool_call_count += 1 - return len(rows), tool_call_count + return len(rows), sum(_tool_calls_len(row[0], scalar=1) for row in rows) def _split_rewind_target(self, target_row: Dict[str, Any], expected_target_content: Any, preserve_compaction_handoff: bool): - """Validate an active rewind target and return its handoff scaffold (or None). - - Raises ``ValueError`` for an inactive / non-user-originated target or a missing - composite carrier, ``RuntimeError`` when the canonical live payload no longer - matches *expected_target_content*. - """ + """Validate an active rewind target; return its handoff scaffold (or None). + ``ValueError``: inactive / non-user-originated target or missing composite + carrier; ``RuntimeError``: live payload no longer matches *expected_target_content*.""" if not target_row.get("active"): raise ValueError("rewind target is not active") from agent.context_compressor import split_user_originated_turn - split_target = target_row.copy() split_target["content"] = self._decode_content(split_target.get("content")) split_target["display_metadata"] = self._decode_display_metadata(split_target.get("display_metadata")) @@ -1454,21 +1237,16 @@ class SessionMessagesMixin: self, session_id: str, target_message_id: int, *, preserve_compaction_handoff: bool = False, expected_active_ids: Optional[List[int]] = None, expected_target_content: Any = None, ) -> Dict[str, Any]: - """Soft-delete (``active=0``) every message with id >= *target_message_id*. - - The target itself goes inactive so the caller can pre-fill it as the next - prompt. Returns ``{"rewound_count", "target_message", "new_head_id"}`` (plus - ``replacement_message_id`` with ``preserve_compaction_handoff``, which archives a - composite summary carrier and inserts its hidden handoff scaffold as the new - head in the same txn). Raises ``ValueError`` when the target is missing or not a - ``user`` row. - - ``expected_active_ids`` / ``expected_target_content`` pin the active row set and - the canonical live payload inside the txn before any mutation (presentation-only - metadata changes do not invalidate a rewind). A live cross-process turn lease - refuses the rewind; expired/dead holders are reclaimed. ``rewind_count`` always - increments, even when the target was already inactive. - """ + """Soft-delete (``active=0``) every message with id >= *target_message_id*, the + target included (the caller pre-fills it as the next prompt). Returns + ``{"rewound_count", "target_message", "new_head_id"}``, plus + ``replacement_message_id`` with ``preserve_compaction_handoff`` (archives a + composite summary carrier, inserts its hidden handoff scaffold as the new head). + ``ValueError`` when the target is missing or not a ``user`` row. + ``expected_active_ids`` / ``expected_target_content`` pin the active set and the + canonical live payload in-txn before any mutation (presentation-only metadata + changes don't invalidate a rewind). A live turn lease refuses the rewind; + expired/dead holders are reclaimed. ``rewind_count`` always increments.""" def _do(conn): self._check_transcript_write_guards( @@ -1478,7 +1256,7 @@ class SessionMessagesMixin: active_rows = conn.execute( "SELECT id FROM messages WHERE session_id = ? AND active = 1 ORDER BY id", (session_id,), ).fetchall() - if [int(active_row[0]) for active_row in active_rows] != expected_active_ids: + if [int(r[0]) for r in active_rows] != expected_active_ids: raise RuntimeError("active transcript changed before the rewind could be persisted") row = conn.execute( "SELECT * FROM messages WHERE id = ? AND session_id = ?", (target_message_id, session_id), @@ -1491,15 +1269,13 @@ class SessionMessagesMixin: f"rewind target must be a 'user' message (got role=" f"{target_row.get('role')!r}, id={target_message_id})" ) - replacement_message_id: Optional[int] = None - replacement: Optional[Dict[str, Any]] = None + replacement_message_id = replacement = None if preserve_compaction_handoff or expected_target_content is not None: replacement = self._split_rewind_target(target_row, expected_target_content, preserve_compaction_handoff) - cursor = conn.execute( + ids = [r[0] for r in conn.execute( "SELECT id FROM messages WHERE session_id = ? AND id >= ? AND active = 1", (session_id, target_message_id), - ) - ids = [r[0] for r in cursor.fetchall()] + ).fetchall()] if ids: conn.execute(f"UPDATE messages SET active = 0 WHERE id IN ({_placeholders(ids)})", ids) if replacement is not None: @@ -1511,13 +1287,11 @@ class SessionMessagesMixin: message_count, tool_call_count = self._active_transcript_counts(conn, session_id) conn.execute( "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", - (message_count, tool_call_count, session_id), - ) + (message_count, tool_call_count, session_id)) head_row = conn.execute( "SELECT MAX(id) FROM messages WHERE session_id = ? AND active = 1", (session_id,), ).fetchone() - new_head_id = head_row[0] if head_row and head_row[0] is not None else None - return target_row, ids, new_head_id, replacement_message_id + return target_row, ids, head_row[0] if head_row else None, replacement_message_id target_row, rewound, new_head_id, replacement_message_id = self._execute_write(_do) # Decode for the prompt-buffer prefill without a second fallible DB operation. @@ -1529,12 +1303,9 @@ class SessionMessagesMixin: def message_count(self, session_id: str = None) -> int: """Count messages, optionally for a specific session.""" - with self._read_ctx() as conn: - if session_id: - cursor = conn.execute("SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,)) - else: - cursor = conn.execute("SELECT COUNT(*) FROM messages") - return cursor.fetchone()[0] + if session_id: + return self._read_one("SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,))[0] + return self._read_one("SELECT COUNT(*) FROM messages")[0] def has_platform_message_id(self, session_id: str, platform_message_id: str) -> bool: """True when a message with *platform_message_id* exists (uses the @@ -1546,13 +1317,10 @@ class SessionMessagesMixin: ) is not None def _is_explicit_fork_child_row(self, session: Dict[str, Any]) -> bool: - """True when *session* is a branch, delegate, or tool child of its parent. - - Markers only count when they point at ``parent_session_id``: compression copies - ``model_config`` onto the continuation, so a delegate's continuation carries - ``_delegate_from=`` and presence-only matching would - misclassify it (same binding as ``_NON_CONTINUATION_CHILD_FILTER_SQL``). - """ + """True when *session* is a branch, delegate, or tool child of its parent. Markers + only count when they point at ``parent_session_id``: compression copies + ``model_config`` onto the continuation, so presence-only matching would misclassify + a delegate's continuation (same binding as ``_NON_CONTINUATION_CHILD_FILTER_SQL``).""" if session.get("source") == "tool": return True raw = session.get("model_config") @@ -1564,12 +1332,9 @@ class SessionMessagesMixin: return False if not isinstance(cfg, dict): return False + markers = (cfg.get("_branched_from"), cfg.get("_delegate_from")) parent_id = session.get("parent_session_id") - branched = cfg.get("_branched_from") - delegated = cfg.get("_delegate_from") - if parent_id: - return branched == parent_id or delegated == parent_id - return branched is not None or delegated is not None + return parent_id in markers if parent_id else any(m is not None for m in markers) def is_explicit_fork_child(self, session_id: str) -> bool: """Public read-only view of :meth:`_is_explicit_fork_child_row`; a missing row @@ -1597,12 +1362,10 @@ class SessionMessagesMixin: return None row = self._read_one( "SELECT generation FROM conversation_generations WHERE source = ? AND session_key = ?", - (source, session_key), - ) - if row is None or row["generation"] is None: + (source, session_key)) + if row is None or row["generation"] is None or int(row["generation"]) <= 0: return None - generation = int(row["generation"]) - return generation if generation > 0 else None + return int(row["generation"]) def clear_messages(self, session_id: str) -> None: """Delete all messages for a session and reset its counters.""" @@ -1614,15 +1377,11 @@ class SessionMessagesMixin: self._execute_write(_do) def purge_stale_tool_call_markers(self, *, dry_run: bool = False, backup: bool = True) -> Dict[str, Any]: - """Permanently clear bare tool-call marker content (e.g. "[memory]") left by - pre-fix sessions. ``_rows_to_conversation`` repairs this in memory on every load, - so this is optional; it just stops the re-scan and removes the bytes. - - Only ``content`` is touched (tool_call pairing unaffected). With ``backup`` a - ``VACUUM INTO`` snapshot (safe against a live connection) is taken first; none - when nothing changes. ``dry_run`` reports without writing or backing up. - Returns ``{"dry_run", "rows_affected", "row_ids", "backup_path"}``. - """ + """Permanently clear bare tool-call marker content (e.g. "[memory]") left by pre-fix + sessions (``_rows_to_conversation`` already repairs it in memory; this stops the + re-scan). Only ``content`` is touched. ``backup``: ``VACUUM INTO`` snapshot first + (none when nothing changes). Returns ``{"dry_run", "rows_affected", "row_ids", + "backup_path"}``.""" from hermes_state import _STALE_TOOL_CALL_MARKER_RE def _find_affected(conn) -> List[int]: @@ -1645,7 +1404,6 @@ class SessionMessagesMixin: backup_path: Optional[str] = None if backup: import datetime - stamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") dest = self.db_path.with_name(f"{self.db_path.name}.pre-clean-markers-backup-{stamp}") with self._lock: @@ -1656,7 +1414,7 @@ class SessionMessagesMixin: def _do(conn): ids = _find_affected(conn) if ids: - conn.execute(f"UPDATE messages SET content = '' WHERE id IN ({','.join('?' * len(ids))})", ids) + conn.execute(f"UPDATE messages SET content = '' WHERE id IN ({_placeholders(ids)})", ids) return ids affected_ids = self._execute_write(_do) diff --git a/hermes_state_titles.py b/hermes_state_titles.py index 9404fa74ad..55c84ee997 100644 --- a/hermes_state_titles.py +++ b/hermes_state_titles.py @@ -79,8 +79,7 @@ class SessionTitlesMixin: nothing overwrites a user name, re-running the titler on an llm row is a no-op). No writer may move a hidden canonical Bot Chat off its title. Read and write are one compare-and-swap transaction, so a manual ``/title`` racing an in-flight - generation is not clobbered. - """ + generation is not clobbered.""" title = self.sanitize_title(title) is_user = source == self.TITLE_SOURCE_USER new_rank = self._title_rank(source) if not is_user else None @@ -166,21 +165,16 @@ class SessionTitlesMixin: if source not in self._TITLE_SOURCE_RANK: raise ValueError(f"invalid title source: {source!r}") return self._write_rowcount( - "UPDATE sessions SET title_source = ? " - "WHERE id = ? AND title IS NOT NULL", + "UPDATE sessions SET title_source = ? WHERE id = ? AND title IS NOT NULL", (source, session_id), ) > 0 def get_session_by_title(self, title: str) -> Optional[Dict[str, Any]]: """Look up a session by exact title. Returns session dict or None.""" row = self._read_one( - "SELECT s.*, " - "COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " - "FROM sessions s " - "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash " - "WHERE s.title = ?", - (title,), - ) + "SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " + "FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash " + "WHERE s.title = ?", (title,)) return self._session_row_dict(row) if row else None def resolve_session_by_title(self, title: str) -> Optional[str]: @@ -191,8 +185,7 @@ class SessionTitlesMixin: numbered = self._read_all( "SELECT id, title, started_at FROM sessions " "WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC", - (f"{_escape_like(title)} #%",), - ) + (f"{_escape_like(title)} #%",)) if numbered: return numbered[0]["id"] return exact["id"] if exact else None @@ -204,13 +197,9 @@ class SessionTitlesMixin: base = match.group(1) if match else base_title rows = self._read_all( "SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'", - (base, f"{_escape_like(base)} #%"), - ) + (base, f"{_escape_like(base)} #%")) if not rows: return base - max_num = 1 # the unnumbered original counts as #1 - for row in rows: - m = re.match(r'^.* #(\d+)$', row["title"]) - if m: - max_num = max(max_num, int(m.group(1))) - return f"{base} #{max_num + 1}" + # The unnumbered original counts as #1. + numbers = [int(m.group(1)) for m in (re.match(r'^.* #(\d+)$', row["title"]) for row in rows) if m] + return f"{base} #{max([1, *numbers]) + 1}" diff --git a/hermes_state_usage.py b/hermes_state_usage.py index f84455053e..75582f89c2 100644 --- a/hermes_state_usage.py +++ b/hermes_state_usage.py @@ -78,6 +78,14 @@ _MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage ( last_seen = excluded.last_seen""" +# Kwargs forwarded verbatim from update_token_counts / record_auxiliary_usage into +# _record_model_usage (the per-route attribution row). +_MODEL_USAGE_FIELDS = frozenset(( + "model", "billing_provider", "billing_base_url", "billing_mode", "input_tokens", "output_tokens", + "cache_read_tokens", "cache_write_tokens", "reasoning_tokens", "estimated_cost_usd", + "actual_cost_usd", "cost_status", "cost_source", "api_call_count")) + + class SessionUsageMixin: """Coalesced token writer, per-model usage rows, billing route.""" @@ -110,16 +118,16 @@ class SessionUsageMixin: to the synchronous path and may raise.""" with self._token_queue_cond: thread = self._token_writer_thread - writer_stopped = self._token_writer_stop and (thread is None or not thread.is_alive()) + writer_alive = thread is not None and thread.is_alive() + writer_stopped = self._token_writer_stop and not writer_alive if not writer_stopped: self._token_queue.append((session_id, kwargs)) - if thread is None or not thread.is_alive(): + if not writer_alive: # Daemon so exit never hangs on accounting; the atexit hook drains # leftovers. ``not is_alive()`` (not ``is None``) respawns a writer # that died from an unexpected escape. thread = threading.Thread( - target=self._token_writer_loop, name="session-db-token-writer", daemon=True, - ) + target=self._token_writer_loop, name="session-db-token-writer", daemon=True) self._token_writer_thread = thread thread.start() if self._token_atexit_hook is None: @@ -245,9 +253,7 @@ class SessionUsageMixin: # Writer stuck mid-apply: leave deltas unapplied rather than race it. logger.warning( "async token accounting: writer did not stop within %.0fs; " - "%d queued delta(s) not persisted", - join_timeout, len(self._token_queue), - ) + "%d queued delta(s) not persisted", join_timeout, len(self._token_queue)) return # Writer gone: apply leftovers synchronously under the same busy protocol. Wait # out a flush caller-drain that already claimed busy — close() nulls the @@ -260,8 +266,7 @@ class SessionUsageMixin: logger.warning( "async token accounting: concurrent drain did not " "finish within %.0fs; %d queued delta(s) not persisted", - join_timeout, len(self._token_queue), - ) + join_timeout, len(self._token_queue)) return self._token_queue_cond.wait(remaining) # busy BEFORE clearing the queue (same ordering as the writer loop). @@ -290,6 +295,7 @@ class SessionUsageMixin: """Update token counters and backfill model if unset. *absolute*=False increments (per-API-call deltas, CLI path); *absolute*=True sets directly (gateway path, where the cached agent holds cumulative totals).""" + usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS} # Ensure the row exists: under concurrent load create_session() may have failed # on locking, and the UPDATE would silently affect 0 rows. self._insert_session_row(session_id, "unknown", model=model) @@ -305,8 +311,7 @@ class SessionUsageMixin: billing_provider if has_accounted_usage else None, billing_base_url if has_accounted_usage else None, billing_mode if has_accounted_usage else None, model if has_accounted_usage else None, - api_call_count, session_id, - ) + api_call_count, session_id) # Per-model attribution: the sessions row keeps one (model, provider) pair, so a # mid-session /model switch would attribute every token to the initial model. # Only the incremental path records here — absolute cumulative updates cannot be @@ -317,18 +322,14 @@ class SessionUsageMixin: row = conn.execute( "SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,), ).fetchone() - existing_model = row["model"] if row is not None else None - existing_provider = row["billing_provider"] if row is not None else None - existing_api_calls = int((row["api_call_count"] if row is not None else 0) or 0) + existing = dict(row) if row is not None else {} # create_session records the requested route before any API call. If that # fails and fallback succeeds, the first accounted usage is the authoritative # route; after that keep the row as is (one row cannot represent mixed usage). first_accounted_route = ( - existing_api_calls == 0 - and has_accounted_usage - and bool(model) + int(existing.get("api_call_count") or 0) == 0 and has_accounted_usage and bool(model) and bool(billing_provider) - and (existing_model != model or existing_provider != billing_provider) + and (existing.get("model") != model or existing.get("billing_provider") != billing_provider) ) if first_accounted_route: conn.execute( @@ -340,51 +341,39 @@ class SessionUsageMixin: ) conn.execute(sql, params) if record_model_usage: - self._record_model_usage( - conn, session_id, model=model, billing_provider=billing_provider, - billing_base_url=billing_base_url, billing_mode=billing_mode, - input_tokens=input_tokens, output_tokens=output_tokens, - cache_read_tokens=cache_read_tokens, cache_write_tokens=cache_write_tokens, - reasoning_tokens=reasoning_tokens, estimated_cost_usd=estimated_cost_usd, - actual_cost_usd=actual_cost_usd, cost_status=cost_status, cost_source=cost_source, - api_call_count=api_call_count, - ) + self._record_model_usage(conn, session_id, **usage) self._execute_write(_do) def _record_model_usage( - self, conn, session_id: str, *, model: Optional[str], billing_provider: Optional[str], - billing_base_url: Optional[str], billing_mode: Optional[str], input_tokens: int, - output_tokens: int, cache_read_tokens: int, cache_write_tokens: int, reasoning_tokens: int, - estimated_cost_usd: Optional[float], actual_cost_usd: Optional[float], - cost_status: Optional[str], cost_source: Optional[str], api_call_count: int, task: str = "", + self, conn, session_id: str, *, model: Optional[str] = None, billing_provider: Optional[str] = None, + billing_base_url: Optional[str] = None, billing_mode: Optional[str] = None, input_tokens: int = 0, + output_tokens: int = 0, cache_read_tokens: int = 0, cache_write_tokens: int = 0, + reasoning_tokens: int = 0, estimated_cost_usd: Optional[float] = None, + actual_cost_usd: Optional[float] = None, cost_status: Optional[str] = None, + cost_source: Optional[str] = None, api_call_count: int = 0, task: str = "", ) -> None: """Accumulate a per-API-call usage delta into session_model_usage, inside the caller's write txn after the ``sessions`` UPDATE. A missing model/provider falls back to the session row (same COALESCE behaviour as the summary update) — except for aux rows (``task`` set), which must NOT inherit the main-loop route (vision - on gemini while the main loop runs anthropic): missing info stays 'unknown'/empty. - """ + on gemini while the main loop runs anthropic): missing info stays 'unknown'/empty.""" row = conn.execute( "SELECT model, billing_provider, billing_base_url, billing_mode " - "FROM sessions WHERE id = ?", - (session_id,), + "FROM sessions WHERE id = ?", (session_id,), ).fetchone() sess = dict(row) if (row is not None and not task) else {} - eff_model = model or sess.get("model") or "unknown" - eff_provider = billing_provider or sess.get("billing_provider") or "" - eff_base_url = billing_base_url or sess.get("billing_base_url") or "" - eff_billing_mode = billing_mode or sess.get("billing_mode") or "" counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)] now = time.time() conn.execute( _MODEL_USAGE_UPSERT_SQL, ( - session_id, eff_model, eff_provider, eff_base_url, eff_billing_mode, task or "", + session_id, model or sess.get("model") or "unknown", + billing_provider or sess.get("billing_provider") or "", + billing_base_url or sess.get("billing_base_url") or "", + billing_mode or sess.get("billing_mode") or "", task or "", api_call_count or 0, *counts, float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0), - cost_status, cost_source, now, now, - ), - ) + cost_status, cost_source, now, now)) def record_auxiliary_usage( self, session_id: str, task: str, *, model: Optional[str] = None, @@ -398,22 +387,13 @@ class SessionUsageMixin: touching the ``sessions`` summary row (the gateway overwrites those counters with absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort: callers must never fail an aux call over accounting.""" + usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS} if not session_id or not task: return + usage["api_call_count"] = 1 if api_call_count is None else int(api_call_count) # FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts. self._insert_session_row(session_id, "unknown") - - def _do(conn): - self._record_model_usage( - conn, session_id, model=model, billing_provider=billing_provider, - billing_base_url=billing_base_url, billing_mode=None, - input_tokens=input_tokens or 0, output_tokens=output_tokens or 0, - cache_read_tokens=cache_read_tokens or 0, cache_write_tokens=cache_write_tokens or 0, - reasoning_tokens=reasoning_tokens or 0, estimated_cost_usd=estimated_cost_usd, - actual_cost_usd=None, cost_status=None, cost_source=None, - api_call_count=1 if api_call_count is None else int(api_call_count), task=task, - ) - self._execute_write(_do) + self._execute_write(lambda conn: self._record_model_usage(conn, session_id, task=task, **usage)) def usage_totals(self, *, min_message_count: int = 1, include_archived: bool = False) -> Dict[str, float]: """Tokens and spend across the whole store (one scan), so the sidebar total does