From 5c3acca66bfb14db631108acab119912c03a991c Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:37:21 -0700 Subject: [PATCH] =?UTF-8?q?refactor(state):=20resume=20=E2=80=94=20verifie?= =?UTF-8?q?d=20messages/compression/titles/usage=20simplification?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- hermes_state_compression.py | 675 ++++------ hermes_state_messages.py | 2419 +++++++++++------------------------ hermes_state_titles.py | 249 ++-- hermes_state_usage.py | 458 +++---- 4 files changed, 1210 insertions(+), 2591 deletions(-) diff --git a/hermes_state_compression.py b/hermes_state_compression.py index cb91823f39..31b4983fd1 100644 --- a/hermes_state_compression.py +++ b/hermes_state_compression.py @@ -16,30 +16,29 @@ from hermes_state_common import _sql_session_last_active, is_automatic_end_reaso # Log-record parity with the origin module (caplog tests pin "hermes_state"). logger = logging.getLogger("hermes_state") +_ENDED_ROW_SQL = "SELECT ended_at, end_reason FROM sessions WHERE id = ?" +_LOCK_ROW_SQL = "SELECT holder, expires_at FROM compression_locks WHERE session_id = ?" +_COOLDOWN_ROW_SQL = ( + "SELECT compression_failure_cooldown_until, compression_failure_error FROM sessions WHERE id = ?" +) + + +def _ended_by_compression(row) -> bool: + return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" + class SessionCompressionMixin: """Compression lineage, cooldown/streak counters, locks and turn leases.""" - def find_live_compression_child( - self, parent_session_id: str - ) -> Optional[Dict[str, Any]]: - """Return the unique live direct child of a compression-ended session. + def find_live_compression_child(self, parent_session_id: str) -> Optional[Dict[str, Any]]: + """The unique live direct child of a compression-ended session, else None. A stale agent whose parent was rotated elsewhere may recover only when the - lineage names exactly one live continuation; more than one fails closed - rather than guessing which transcript owns later messages.""" + lineage names exactly one live continuation; more than one fails closed.""" if not parent_session_id: return None with self._read_ctx() as conn: - parent = conn.execute( - "SELECT ended_at, end_reason FROM sessions WHERE id = ?", - (parent_session_id,), - ).fetchone() - if ( - parent is None - or parent["ended_at"] is None - or parent["end_reason"] != "compression" - ): + if not _ended_by_compression(conn.execute(_ENDED_ROW_SQL, (parent_session_id,)).fetchone()): return None rows = conn.execute( """ @@ -61,28 +60,16 @@ class SessionCompressionMixin: return self._session_row_dict(rows[0]) if len(rows) == 1 else None def reopen_orphaned_compression_session(self, session_id: str) -> bool: - """Reopen a compression parent only when no continuation was published. - - Publication is atomic now, but older builds could leave a closed parent - after an interrupted handoff. Conservative by design: an active lease or - any canonical child means another path owns the lineage — fail closed.""" + """Reopen a compression parent only when no continuation was published (older + builds could leave a closed parent after an interrupted handoff). Conservative: + an active lease or any canonical child means another path owns the lineage.""" if not session_id: return False def _do(conn): - parent = conn.execute( - "SELECT ended_at, end_reason FROM sessions WHERE id = ?", - (session_id,), - ).fetchone() - if ( - parent is None - or parent["ended_at"] is None - or parent["end_reason"] != "compression" - ): + if not _ended_by_compression(conn.execute(_ENDED_ROW_SQL, (session_id,)).fetchone()): return False - - # Any non-branch/non-delegate/non-tool child is a continuation, ended - # or not; reopening past it could give one lineage a second live head. + # Any non-branch/non-delegate/non-tool child is a continuation, ended or not. child = conn.execute( """ SELECT 1 @@ -97,17 +84,11 @@ class SessionCompressionMixin: ).fetchone() if child is not None: return False - - # refresh_compression_lock() lets an owner revive its own expired - # row, so reclaim it inside this write transaction before reopening: - # refresh-first makes the lease active and aborts recovery; - # recovery-first deletes the holder so a later refresh can't resurrect it. + # refresh_compression_lock() lets an owner revive its own expired row, so + # reclaim it inside this write txn: refresh-first makes the lease active and + # aborts recovery; recovery-first deletes the holder so a refresh can't resurrect it. now = time.time() - lock_row = conn.execute( - "SELECT holder, expires_at FROM compression_locks " - "WHERE session_id = ?", - (session_id,), - ).fetchone() + lock_row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone() if lock_row is not None: expires_at = lock_row["expires_at"] if expires_at is None or float(expires_at) >= now: @@ -119,20 +100,56 @@ class SessionCompressionMixin: ) 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,), ) - # rowcount==1 is guaranteed by the parent SELECT in this same BEGIN - # IMMEDIATE transaction. If a False return is ever added past this - # point, raise instead: _execute_write commits the lease DELETE above unless _do raises. + # 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. return updated.rowcount == 1 return bool(self._execute_write(_do)) + 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*.""" + system_prompt_hash = self._store_system_prompt(conn, system_prompt) + conn.execute( + """INSERT INTO sessions ( + id, source, model, model_config, system_prompt, + system_prompt_hash, + parent_session_id, cwd, git_branch, git_repo_root, + profile_name, user_id, session_key, chat_id, chat_type, + 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"], + 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(), + ), + ) + def publish_compression_child( self, *, @@ -154,35 +171,29 @@ class SessionCompressionMixin: ) -> None: """Atomically close a parent and publish its durable compression child. - Closure, child row, and handoff commit in one transaction: readers see - the live parent or a complete child, never an ended parent with a - missing/empty child. + Closure, child row, and handoff commit in one transaction: readers see the live + parent or a complete child, never an ended parent with a missing/empty child. - *watermark* (parent's ``get_active_message_watermark`` at compression - start): parent rows with ``id > watermark`` — appends landed during the - slow summary call — are column-cloned into the child AFTER the handoff - so they survive rotation. *watermark_ceiling* bounds the clone: the - rotation path flushes its OWN transcript to the parent just before - publishing and those rows are already in the handoff, so the caller - captures ``MAX(id)`` right BEFORE that flush and only - ``(watermark, watermark_ceiling]`` is foreign tail. ``None`` = unbounded. + *watermark* (parent's ``get_active_message_watermark`` at compression start): + parent rows with ``id > watermark`` — appends landed during the slow summary — + are column-cloned into the child AFTER the handoff. *watermark_ceiling* bounds + the clone: the rotation path flushes its OWN transcript to the parent just + before publishing and those rows are already in the handoff, so only + ``(watermark, watermark_ceiling]`` is foreign tail (``None`` = unbounded). - *require_lease_refresh* + *compression_lock_holder* refreshes the lease - on the same ``conn`` before the expiry check (no TOCTOU window), so a - refresher that died on transient DB errors gets one last chance.""" + *require_lease_refresh* + *compression_lock_holder* refreshes the lease on the + same ``conn`` before the expiry check (no TOCTOU window), so a refresher that + died on transient DB errors gets one last chance.""" from hermes_state import CompressionSessionBusyError + def _do(conn): if require_lease_refresh and compression_lock_holder: 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( - "SELECT holder, expires_at FROM compression_locks WHERE session_id = ?", - (parent_session_id,), - ).fetchone() + lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone() if require_compression_lease and ( lock_row is None or not compression_lock_holder @@ -202,14 +213,11 @@ class SessionCompressionMixin: if parent is None: raise RuntimeError(f"Compression parent not found: {parent_session_id}") if parent["ended_at"] is not None: - # An ended stamp from AUTOMATIC cleanup (tui_shutdown, ws_disconnect, - # orphan reap, idle/LRU evict) is stale by construction — this lease - # holder is still continuing the conversation. Left alone it wedges - # rotation forever (every attempt aborts here; each pre-publish flush - # re-grows the parent until the provider rejects it). Clear it; the - # closure UPDATE below re-stamps end_reason='compression'. Deliberate - # boundaries (compression, session_reset, explicit close) still fail - # closed — another path owns the lineage. + # An AUTOMATIC end stamp (tui_shutdown, ws_disconnect, orphan reap, + # idle/LRU evict) is stale by construction — this lease holder is still + # continuing the conversation, and left alone it wedges rotation forever. + # Clear it; the closure UPDATE below re-stamps end_reason='compression'. + # Deliberate boundaries still fail closed. if is_automatic_end_reason(parent["end_reason"]): conn.execute( "UPDATE sessions SET ended_at = NULL, end_reason = NULL " @@ -217,90 +225,34 @@ class SessionCompressionMixin: (parent_session_id,), ) else: - raise RuntimeError( - f"Compression parent already ended: {parent_session_id}" - ) + raise RuntimeError(f"Compression parent already ended: {parent_session_id}") if not messages: raise RuntimeError("Compression child handoff must not be empty") - system_prompt_hash = self._store_system_prompt(conn, system_prompt) - - conn.execute( - """INSERT INTO sessions ( - id, source, model, model_config, system_prompt, - system_prompt_hash, - parent_session_id, cwd, git_branch, git_repo_root, - profile_name, user_id, session_key, chat_id, chat_type, - 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"], - parent["git_repo_root"], - # Same contract as _insert_session_row's compression-fork backfill: - # child stays on the parent's profile and keeps gateway routing/ - # origin columns so peer recovery works after a boundary crash. No - # owner on either side (legacy NULL parent) → stamp this store's - # profile so the child doesn't extend the unowned lineage. - 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(), - ), - ) - total_messages, total_tool_calls = self._insert_message_rows( - conn, child_session_id, messages + 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, ) + 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 (see docstring) into the - # child after the handoff: column-exact except id/session_id; + # 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)) - tail_rows = conn.execute( + 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 > ?" f"{_ceiling_clause} ORDER BY id", _params, - ).fetchall() - if tail_rows: - tail_ids = [int(r["id"]) for r in tail_rows] - placeholders = ",".join("?" for _ in tail_ids) - clone_cols = [ - c for c in self._message_column_names(conn) - if c not in ("id", "session_id", "active", "compacted") - ] - col_list = ", ".join(clone_cols) - 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", - [child_session_id, *tail_ids], - ) + ) + if tail_ids: + self._clone_message_rows(conn, tail_ids, session_id=child_session_id) total_messages += len(tail_ids) - for r in tail_rows: - raw = r["tool_calls"] - if raw: - try: - parsed = json.loads(raw) if isinstance(raw, str) else raw - total_tool_calls += len(parsed) if isinstance(parsed, list) else 0 - except (TypeError, ValueError): - pass + 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), @@ -311,111 +263,66 @@ class SessionCompressionMixin: (time.time(), parent_session_id), ) if updated.rowcount != 1: - raise RuntimeError( - f"Compression parent changed during publication: {parent_session_id}" - ) + raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}") self._execute_write(_do) - def record_compression_failure_cooldown( - self, - session_id: str, - cooldown_until: float, - error: Optional[str] = None, - ) -> None: - """Persist the active compression-failure cooldown for a session.""" + def _write_sql_logged(self, op: str, session_id: str, sql: str, params) -> None: + """``_write_sql`` that logs (never raises) on ``sqlite3.Error``.""" + try: + self._write_sql(sql, params) + except sqlite3.Error as exc: + logger.warning("%s(%s) failed: %s", op, session_id, exc) + + def record_compression_failure_cooldown(self, session_id: str, cooldown_until: float, error: Optional[str] = None) -> None: + """Persist the active compression-failure cooldown. Merge-max with any longer + live deadline so a later shorter write can't reopen the thrash window; error + always takes the latest diagnostic.""" if not session_id: return + self._write_sql_logged( + "record_compression_failure_cooldown", session_id, + "UPDATE sessions SET compression_failure_cooldown_until = CASE " + "WHEN compression_failure_cooldown_until IS NOT NULL " + " AND compression_failure_cooldown_until > ? " + "THEN compression_failure_cooldown_until ELSE ? END, " + "compression_failure_error = ? WHERE id = ?", + (cooldown_until, cooldown_until, error, session_id), + ) - try: - # Merge-max with any longer live deadline so a later shorter write - # can't reopen the thrash window; error always takes the latest diagnostic. - self._write_sql( - "UPDATE sessions SET compression_failure_cooldown_until = CASE " - "WHEN compression_failure_cooldown_until IS NOT NULL " - " AND compression_failure_cooldown_until > ? " - "THEN compression_failure_cooldown_until ELSE ? END, " - "compression_failure_error = ? WHERE id = ?", - (cooldown_until, cooldown_until, error, session_id), - ) - except sqlite3.Error as exc: - logger.warning( - "record_compression_failure_cooldown(%s) failed: %s", - session_id, exc, - ) - - def get_compression_failure_cooldown( - self, - session_id: str, - ) -> Optional[Dict[str, Any]]: - """Return the active compression-failure cooldown for ``session_id``.""" + def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]: + """Return the active (unexpired) compression-failure cooldown, or None.""" if not session_id: return None now = time.time() - row = self._read_one( - "SELECT compression_failure_cooldown_until, compression_failure_error " - "FROM sessions WHERE id = ?", - (session_id,), - ) - if row is None: + row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) + if row is None or row[0] is None: return None - cooldown_until = row[0] - if cooldown_until is None: - return None - cooldown_until = float(cooldown_until) + cooldown_until = float(row[0]) if cooldown_until <= now: return None - error = row[1] - return { - "cooldown_until": cooldown_until, - "remaining_seconds": cooldown_until - now, - "error": error, - } + return {"cooldown_until": cooldown_until, "remaining_seconds": cooldown_until - now, "error": row[1]} - def get_compression_failure_cooldown_row( - self, - session_id: str, - ) -> Dict[str, Any]: - """Exact stored cooldown columns, no expiry filtering. Compression - cancellation uses this under its session lease so rollback preserves an - expired, partially-null, or absent row exactly instead of coercing it - through the active-cooldown API.""" - if not session_id: - return {"session_exists": False, "cooldown_until": None, "error": None} - row = self._read_one( - "SELECT compression_failure_cooldown_until, compression_failure_error " - "FROM sessions WHERE id = ?", - (session_id,), - ) + 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} - cooldown_until = row[0] - error = row[1] return { "session_exists": True, - "cooldown_until": ( - float(cooldown_until) if cooldown_until is not None else None - ), - "error": error, + "cooldown_until": float(row[0]) if row[0] is not None else None, + "error": 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 rollback API propagates write and verification failures: cancellation - must not be reported mutation-free when compensation failed.""" - expected_exists = bool(snapshot.get("session_exists", False)) - if not expected_exists: - actual = self.get_compression_failure_cooldown_row(session_id) - if actual.get("session_exists", False): - raise RuntimeError( - "cannot restore absent compression cooldown row: session now exists" - ) + 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 + rollback API propagates write and verification failures: cancellation must not + be reported mutation-free when compensation failed.""" + if not snapshot.get("session_exists", False): + if self.get_compression_failure_cooldown_row(session_id).get("session_exists", False): + raise RuntimeError("cannot restore absent compression cooldown row: session now exists") return - deadline = snapshot.get("cooldown_until") error = snapshot.get("error") @@ -426,9 +333,7 @@ class SessionCompressionMixin: (deadline, error, session_id), ) if cursor.rowcount != 1: - raise RuntimeError( - f"compression cooldown rollback session missing: {session_id}" - ) + raise RuntimeError(f"compression cooldown rollback session missing: {session_id}") self._execute_write(_do) actual = self.get_compression_failure_cooldown_row(session_id) @@ -447,27 +352,19 @@ class SessionCompressionMixin: """Clear any persisted compression-failure cooldown for a session.""" if not session_id: return - - try: - self._write_sql( - "UPDATE sessions SET compression_failure_cooldown_until = NULL, " - "compression_failure_error = NULL WHERE id = ?", - (session_id,), - ) - except sqlite3.Error as exc: - logger.warning( - "clear_compression_failure_cooldown(%s) failed: %s", - session_id, exc, - ) + 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,), + ) 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, NULL, or unparsable value also reads as ``zero``.""" + """Read one numeric ``sessions`` column clamped at ``zero``; a missing session, + NULL, or unparsable value also reads as ``zero``.""" if not session_id: return zero - row = self._read_one( - f"SELECT {column} FROM sessions WHERE id = ?", (session_id,) - ) + row = self._read_one(f"SELECT {column} FROM sessions WHERE id = ?", (session_id,)) if row is None: return zero try: @@ -488,9 +385,9 @@ class SessionCompressionMixin: ) def get_compression_ineffective_count(self, session_id: str) -> int: - """Persisted ineffective-compaction strike count: the durable half of - the built-in compressor's anti-thrash guard, so a fresh compressor bound - to a resumed session inherits an armed/tripped guard across restarts.""" + """Persisted ineffective-compaction strike count — the durable half of the + built-in compressor's anti-thrash guard, so a fresh compressor bound to a resumed + session inherits an armed/tripped guard across restarts.""" return self._read_session_number("compression_ineffective_count", session_id, int, 0) def set_compression_ineffective_count(self, session_id: str, count: int) -> None: @@ -502,10 +399,8 @@ class SessionCompressionMixin: ) def get_compression_recovery_deadline(self, session_id: str) -> float: - """Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). - Durable because the gateway rebuilds the compressor every turn / cache - eviction: a process-local deadline restarted on each rebuild, so a - tripped session never earned its probe.""" + """Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable + because the gateway rebuilds the compressor every turn / cache eviction.""" return self._read_session_number("compression_recovery_deadline", session_id, float, 0.0) def set_compression_recovery_deadline(self, session_id: str, deadline: float) -> None: @@ -516,39 +411,23 @@ class SessionCompressionMixin: normalized = max(0.0, float(deadline or 0.0)) except (TypeError, ValueError): normalized = 0.0 - stored = normalized if normalized > 0.0 else None - self._write_sql( "UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?", - (stored, session_id), + (normalized if normalized > 0.0 else None, session_id), ) - def refresh_compression_lock( - self, - session_id: str, - holder: str, - ttl_seconds: float = 300.0, - ) -> bool: + 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. - Ownership is decided by ``holder`` alone, deliberately NOT ``expires_at``: - a live owner whose refresher stalled past its TTL (GC pause, loaded CI - runner, slow write escaping ``_execute_write``'s retry budget) must be - able to revive its still-unclaimed row. Requiring ``expires_at >= now`` - made such a stall permanent — every later refresh matched 0 rows and the - owner kept compressing/rotating with no lease, exactly the window in - which a competing path can fork the lineage. - - It cannot resurrect a lock someone else took: SQLite serialises writes, - so :meth:`try_acquire_compression_lock`'s reclaim (DELETE-expired + - INSERT-or-IGNORE) never interleaves with this UPDATE. Reclaim-first - replaces ``holder`` and this matches nothing; refresh-first pushes - ``expires_at`` forward and the reclaimer's DELETE matches nothing.""" + Ownership is decided by ``holder`` alone, deliberately NOT ``expires_at``: a live + owner whose refresher stalled past its TTL must be able to revive its still- + unclaimed row, otherwise it keeps compressing with no lease — the window in + which a competing path can fork the lineage. It cannot resurrect a lock someone + else took: SQLite serialises writes, so the reclaim (DELETE-expired + INSERT OR + IGNORE) never interleaves with this UPDATE.""" if not session_id or not holder: return False - now = time.time() - expires_at = now + ttl_seconds - + expires_at = time.time() + ttl_seconds try: return self._write_rowcount( "UPDATE compression_locks SET expires_at = ? " @@ -556,27 +435,17 @@ class SessionCompressionMixin: (expires_at, session_id, holder), ) > 0 except sqlite3.Error as exc: - logger.warning( - "refresh_compression_lock(%s) failed: %s", - session_id, exc, - ) + logger.warning("refresh_compression_lock(%s) failed: %s", session_id, exc) return False - def try_acquire_compression_lock( - self, - session_id: str, - holder: str, - ttl_seconds: float = 300.0, - ) -> bool: + def try_acquire_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool: """Try to atomically acquire the compression lock for ``session_id``. - ``True``: caller owns the lock and must :meth:`release_compression_lock`. - ``False``: another holder owns a live lock and the caller MUST NOT - compress — its rotation would race the holder's and split the lineage. - Expired locks and structured holders whose local ``pid=`` is dead are - reclaimed transparently, so a gateway killed mid-compression doesn't - stall its replacement for the full TTL. Single-transaction DELETE-expired - + INSERT-or-IGNORE + SELECT-to-confirm; SQLite serialises writes, so it's atomic.""" + ``False``: another holder owns a live lock and the caller MUST NOT compress (its + rotation would split the lineage). Expired locks and structured holders whose + local ``pid=`` is dead are reclaimed transparently. Single-transaction DELETE- + expired + INSERT OR IGNORE + SELECT-to-confirm (INSERT OR IGNORE gives no + rowcount signal).""" from hermes_state import _compression_lock_holder_process_is_dead if not session_id: return False @@ -585,90 +454,54 @@ class SessionCompressionMixin: def _do(conn): reclaimed_holder = None - row = conn.execute( - "SELECT holder, expires_at FROM compression_locks " - "WHERE session_id = ?", - (session_id,), - ).fetchone() + row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone() if row is not None: - current_holder = ( - row[0] - ) - current_expires_at = ( - row[1] - ) - if ( - current_expires_at < now - or _compression_lock_holder_process_is_dead(current_holder) - ): + 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 - # INSERT OR IGNORE gives no rowcount signal — verify ownership via SELECT. 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() - acquired = row is not None and ( - row[0] - ) == holder - return acquired, reclaimed_holder + 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 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, + "Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder, ) return bool(acquired) except sqlite3.Error as exc: - logger.warning( - "try_acquire_compression_lock(%s) failed: %s", - session_id, exc, - ) - # False makes the caller skip compression — the safe behaviour - # when the lock subsystem is broken. + # False makes the caller skip compression — safe when the lock subsystem is broken. + logger.warning("try_acquire_compression_lock(%s) failed: %s", session_id, exc) return False def release_compression_lock(self, session_id: str, holder: str) -> None: - """Release the compression lock for ``session_id`` iff we own it. Idempotent - when the lock is gone or reclaimed; the ``holder`` check stops a late - compressor clobbering someone else's fresh lock.""" + """Release the compression lock iff we own it; idempotent when gone/reclaimed.""" if not session_id: return - - try: - self._write_sql( - "DELETE FROM compression_locks " - "WHERE session_id = ? AND holder = ?", - (session_id, holder), - ) - except sqlite3.Error as exc: - logger.warning( - "release_compression_lock(%s) failed: %s", - session_id, exc, - ) + self._write_sql_logged( + "release_compression_lock", session_id, + "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. - Must share the connection of the lease INSERT/UPDATE/DELETE: a failed - ``get_session`` must not yield a child id the write then persists - (refresh would walk to the parent and fail-close). Markers bind to - ``parent_session_id`` (as in ``_NON_CONTINUATION_CHILD_FILTER_SQL``). - Lock errors propagate so ``_execute_write`` / ``acquire_session_turn_lease`` can retry.""" + Must share the connection of the lease INSERT/UPDATE/DELETE: a failed lookup + must not yield a child id the write then persists. Markers bind to + ``parent_session_id``. Lock errors propagate so ``_execute_write`` can retry.""" if not session_id: return session_id @@ -684,11 +517,7 @@ class SessionCompressionMixin: seen = {session_id} while current: parent_id = current.get("parent_session_id") - if ( - not parent_id - or parent_id in seen - or self._is_explicit_fork_child_row(current) - ): + if not parent_id or parent_id in seen or self._is_explicit_fork_child_row(current): break parent = _row(parent_id) if not parent or parent.get("end_reason") != "compression": @@ -698,29 +527,19 @@ class SessionCompressionMixin: return str(current.get("id") or session_id) if current else session_id def _session_turn_lease_key(self, session_id: str) -> str: - """Return the stable serialization key for every compression segment. - - Acquire/refresh/release resolve this inside their write transaction; this - is for tests/diagnostics. It does not swallow lock errors — a swallowed - walk plus a later successful write was the fail-open that replayed the - post-rotation refresh miss.""" + """Stable serialization key for every compression segment (tests/diagnostics; + the write paths resolve it inside their own txn). Does not swallow lock errors.""" if not session_id: return session_id with self._read_ctx() as conn: return self._session_turn_lease_key_on_conn(conn, session_id) def try_acquire_session_turn_lease( - self, - session_id: str, - holder: str, - *, - ttl_seconds: float = 300.0, - patience_s: Optional[float] = None, + self, session_id: str, holder: str, *, ttl_seconds: float = 300.0, patience_s: Optional[float] = None, ) -> bool: - """Atomically acquire the cross-process turn lease for a conversation. - Compression rotates a session into child segments, so the durable key is - the lineage root, not the current segment id. The walk, the INSERT, and - reclaim of expired or dead-local-PID leases share one write transaction.""" + """Atomically acquire the cross-process turn lease for a conversation (keyed by + the lineage root). The walk, the INSERT, and reclaim of expired or dead-local-PID + leases share one write transaction.""" from hermes_state import _compression_lock_holder_process_is_dead if not session_id or not holder: return False @@ -736,10 +555,7 @@ class SessionCompressionMixin: ).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) - ): + 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 = ?", @@ -752,8 +568,7 @@ class SessionCompressionMixin: (conversation_id, holder, now, expires_at), ) owner = conn.execute( - "SELECT holder FROM session_turn_leases WHERE conversation_id = ?", - (conversation_id,), + "SELECT holder FROM session_turn_leases WHERE conversation_id = ?", (conversation_id,), ).fetchone() return owner is not None and owner["holder"] == holder @@ -774,10 +589,9 @@ class SessionCompressionMixin: ) -> bool: """Wait for a cross-process turn lease without holding a SQLite lock. - ``on_wait(elapsed)`` is best-effort: called when the first attempt fails - (elapsed ~0) and about every ``wait_notice_interval_seconds`` after, so - UIs can show another process holds the conversation. ``should_abort()`` - True (e.g. ``/stop``) returns False at once, not after ``wait_seconds``.""" + ``on_wait(elapsed)`` is best-effort: called when the first attempt fails and + about every ``wait_notice_interval_seconds`` after. ``should_abort()`` True + (e.g. ``/stop``) returns False at once.""" from hermes_state import classify_persistence_error deadline = time.monotonic() + max(0.0, float(wait_seconds)) wait_started = None @@ -789,22 +603,15 @@ class SessionCompressionMixin: if should_abort(): return False except Exception: - logger.debug( - "session turn lease should_abort callback failed", - exc_info=True, - ) + 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 (compression publish, large flushes) - # can exhaust one write-patience budget; keep polling until - # wait_seconds or should_abort. + # Long holder transactions can exhaust one write-patience budget; keep + # polling until wait_seconds or should_abort. if classify_persistence_error(exc) != "locked": raise now = time.monotonic() @@ -814,27 +621,16 @@ class SessionCompressionMixin: if wait_started is None: wait_started = now if on_wait is not None and ( - last_notice_at is None - or notice_every == 0.0 - or (now - last_notice_at) >= notice_every + last_notice_at is None or notice_every == 0.0 or (now - last_notice_at) >= notice_every ): try: on_wait(max(0.0, now - wait_started)) except Exception: - logger.debug( - "session turn lease on_wait callback failed", - exc_info=True, - ) + logger.debug("session turn lease on_wait callback failed", exc_info=True) last_notice_at = now time.sleep(min(max(0.01, float(poll_interval_seconds)), remaining)) - def refresh_session_turn_lease( - self, - session_id: str, - holder: str, - *, - ttl_seconds: float = 300.0, - ) -> bool: + def refresh_session_turn_lease(self, session_id: str, holder: str, *, ttl_seconds: float = 300.0) -> bool: """Extend a turn lease only while ``holder`` still owns it.""" if not session_id or not holder: return False @@ -867,24 +663,20 @@ class SessionCompressionMixin: self._execute_write(_do) def get_compression_lock_holder(self, session_id: str) -> Optional[str]: - """Return the current (non-expired) holder for ``session_id``, or None. - Diagnostic only — not part of the locking protocol.""" + """Current (non-expired) holder for ``session_id``, or None. Diagnostic only.""" if not session_id: return None - now = time.time() row = self._read_one( "SELECT holder FROM compression_locks " "WHERE session_id = ? AND expires_at >= ?", - (session_id, now), + (session_id, time.time()), ) - if row is None: - return None - return row[0] + return None if row is None else row[0] def finalize_orphaned_compression_sessions(self) -> int: - """Mark orphaned compression continuations (parent ended by compression; - child has messages, no end_reason/ended_at, api_call_count=0) as - ``orphaned_compression``. Non-destructive: messages are preserved.""" + """Mark orphaned compression continuations (parent ended by compression; child + 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): @@ -917,29 +709,22 @@ class SessionCompressionMixin: return self._execute_write(_do) or 0 def get_compression_chain(self, session_id: str) -> List[str]: - """Walk the compression-continuation chain forward and return every id. + """Walk the compression-continuation chain forward: root-first through the tip + (``[session_id]`` when no continuation). ``get_compression_tip`` is this walk's + last element. - Root-first, ending at the tip; ``[session_id]`` when no continuation - exists. ``get_compression_tip`` is this walk's last element — one - implementation so the two can never disagree. - - A continuation is a child of a session with ``end_reason='compression'``. - Older builds also required ``child.started_at >= parent.ended_at``; - too brittle — gateway + compression races can insert the real - continuation before the parent's ``ended_at`` is written while a stale - websocket later creates a sibling that passes the timestamp test, so - desktop resume followed the sibling and recent messages looked "lost". - Instead: follow only children of compression-ended parents, exclude - explicit branch/delegate/tool children, and prefer children that continue - the chain (``end_reason='compression'``) or are still live over stale - closed siblings such as ``ws_orphan_reap``.""" + A continuation is a child of a session with ``end_reason='compression'``. The + old ``child.started_at >= parent.ended_at`` test was too brittle (gateway + + compression races insert the real continuation before ``ended_at`` is written, + while a stale websocket later creates a sibling that passes it). Instead exclude + branch/delegate/tool children and prefer children that continue the chain or + 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() - # Defensive bound; chains this deep are pathological. - for _ in range(100): + for _ in range(100): # defensive bound; chains this deep are pathological with self._read_ctx() as conn: - cursor = conn.execute( + row = conn.execute( f""" SELECT child.id FROM sessions parent @@ -961,8 +746,7 @@ class SessionCompressionMixin: LIMIT 1 """, (current,), - ) - row = cursor.fetchone() + ).fetchone() if row is None: return chain child_id = row["id"] @@ -974,8 +758,8 @@ class SessionCompressionMixin: return chain def get_compression_tip(self, session_id: str) -> Optional[str]: - """Live tip of a compression chain (walk semantics: ``get_compression_chain``); - the input id when no continuation exists.""" + """Live tip of a compression chain (``get_compression_chain`` semantics); the + input id when no continuation exists.""" chain = self.get_compression_chain(session_id) return chain[-1] if chain else session_id @@ -991,7 +775,6 @@ class SessionCompressionMixin: session = self.get_session(session_id) if not session or self._is_explicit_fork_child_row(session): return [session_id] if session else [] - root = session ancestors = {root["id"]} while self._is_compression_child_row(root): @@ -1000,7 +783,6 @@ class SessionCompressionMixin: break root = parent ancestors.add(root["id"]) - lineage = [root["id"]] seen = {root["id"]} current = root @@ -1013,18 +795,11 @@ class SessionCompressionMixin: """, (current["id"],), ) - next_child = None - for row in rows: - candidate = dict(row) - if self._is_compression_child_row(candidate): - next_child = candidate - break + next_child = next((dict(row) for row in rows if self._is_compression_child_row(dict(row))), None) if not next_child or next_child["id"] in seen: break lineage.append(next_child["id"]) seen.add(next_child["id"]) current = next_child - if current["id"] == session_id: - # Later tips are included only when the requested session itself was compacted. - continue + # Later tips are included only when the requested session itself was compacted. return lineage if session_id in lineage else [session_id] diff --git a/hermes_state_messages.py b/hermes_state_messages.py index 4e5091d46c..01e561c829 100644 --- a/hermes_state_messages.py +++ b/hermes_state_messages.py @@ -1,10 +1,9 @@ """Transcript persistence for SessionDB. -Mixin split out of ``hermes_state.py``; bound onto ``SessionDB`` via the MRO -and built on its ``_read_ctx`` / ``_execute_write`` / ``_write_sql`` / -``_read_one`` / ``_read_all`` primitives. Covers message append / replace / -rewind, reactions, resume-conversation assembly and replayed-user-message -duplicate detection. +Mixin bound onto ``SessionDB`` via the MRO, built on its ``_read_ctx`` / +``_execute_write`` / ``_write_sql`` / ``_read_one`` / ``_read_all`` primitives. +Covers message append / replace / rewind, reactions, resume-conversation +assembly and replayed-user-message duplicate detection. """ from __future__ import annotations @@ -17,42 +16,96 @@ from typing import Any, Dict, List, Optional, Tuple from agent.context_compressor import _DB_PERSISTED_MARKER as _DB_PERSISTED_MARKER_KEY from agent.memory_manager import sanitize_context from agent.message_sanitization import _sanitize_surrogates -from hermes_state_common import ( - _RESET_END_REASONS, - _RESET_END_REASONS_SQL, - _legacy_reset_child_sql, -) +from hermes_state_common import _RESET_END_REASONS, _RESET_END_REASONS_SQL, _legacy_reset_child_sql # Log-record parity with the origin module (caplog tests pin "hermes_state"). logger = logging.getLogger("hermes_state") +# One INSERT shape for every message writer (append, batch, replace, compact, import). +_INSERT_MESSAGE_SQL = """INSERT INTO messages (session_id, role, content, tool_call_id, + tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason, + reasoning, reasoning_content, reasoning_details, codex_reasoning_items, + codex_message_items, platform_message_id, observed, _compressed_summary, active, api_content, display_kind, display_metadata) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""" + +_ENDED_BY_COMPRESSION_SQL = "SELECT ended_at, end_reason FROM sessions WHERE id = ?" +_COMPRESSION_LOCK_ROW_SQL = "SELECT holder, expires_at FROM compression_locks WHERE session_id = ?" +_TURN_LEASE_ROW_SQL = "SELECT holder, expires_at FROM session_turn_leases WHERE conversation_id = ?" +_DISPLAY_ACTIVE_CLAUSE = " AND (active = 1 OR compacted = 1)" +_DELETE_COMPRESSION_LOCK_SQL = "DELETE FROM compression_locks WHERE session_id = ? AND holder = ?" + + +def _placeholders(items) -> str: + return ",".join("?" for _ in items) + + +def _json_or(raw: Any, fallback: Any, warning: str) -> Any: + """``json.loads(raw)``; on failure log *warning* and return *fallback*.""" + try: + return json.loads(raw) + except (json.JSONDecodeError, TypeError): + logger.warning(warning) + return fallback + + +def _tool_calls_len(raw: Any) -> int: + """Tool-call count of a stored ``tool_calls`` column (0 unless a JSON list).""" + 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 + + +def _coerce_timestamp(value: Any, default: float) -> float: + """Explicit message timestamp (datetime or number) or *default* when invalid.""" + if value is None: + return default + try: + return float(value.timestamp()) if hasattr(value, "timestamp") else float(value) + except (TypeError, ValueError): + logger.debug("Ignoring invalid explicit message timestamp: %r", value) + return default + + +def _parse_tool_calls(tool_calls: Any) -> Any: + """tool_calls may be a list (live agent) or a JSON string (import/export); parse + first so json.dumps never double-encodes.""" + if isinstance(tool_calls, str): + try: + return json.loads(tool_calls) + except (json.JSONDecodeError, TypeError): + return [] + return tool_calls + + +def _tool_calls_count(tool_calls: Any) -> int: + if tool_calls is None: + return 0 + return len(tool_calls) if isinstance(tool_calls, list) else 1 + + +def _ended_by_compression(row) -> bool: + return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" + 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. + """Advance this peer's conversation generation past a boundary, inside the + transaction that writes the boundary. - Called inside the transaction that writes the boundary, so the - generation and the ``end_reason`` that caused it commit together. - - Only ``_RESET_END_REASONS`` count: ``compression`` continues one - conversation, and an accidental close is not a replacement. Rows with - no ``session_key`` have no routing peer to advance. - - The counter deliberately does NOT read the session rows. An aggregate - over them (COUNT/MAX of boundaries) can return a pair it already - emitted once ``delete_session()`` or bulk pruning removes an ended row, - which would hand a new conversation a retired affinity identity. This - value only ever increments, so a generation is never reused for a peer - even if every row behind it is gone. + 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. """ if end_reason not in _RESET_END_REASONS: return - row = conn.execute( - "SELECT source, session_key FROM sessions WHERE id = ?", - (session_id,), - ).fetchone() + row = conn.execute("SELECT source, session_key FROM sessions WHERE id = ?", (session_id,)).fetchone() if row is None: return source = str(row["source"] or "").strip() @@ -71,62 +124,34 @@ class SessionMessagesMixin: @classmethod def _encode_content(cls, content: Any) -> Any: - """Serialize structured (list/dict) message content for sqlite. - - sqlite3 can only bind ``str``, ``bytes``, ``int``, ``float``, and ``None`` - to query parameters. Multimodal messages have ``content`` as a list of - parts (``[{"type": "text", ...}, {"type": "image_url", ...}]``), which - raises ``ProgrammingError: Error binding parameter N: type 'list' is - not supported`` when bound directly. - - Returns the value unchanged when it's already a safe scalar, or a - sentinel-prefixed JSON string for lists/dicts. Paired with - :meth:`_decode_content` on read. + """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`. """ if isinstance(content, str): - # Lone UTF-16 surrogates reach here inside tool results scraped - # from the web/social platforms (the same input that crashed the - # guardrail hasher). The proactive sanitizer upstream only cleans - # the *api_messages* copy, and the recovery sanitizer only runs - # after the API call itself raises — which it no longer does — so - # the canonical history keeps them and this write is where they - # land. Left raw, sqlite3 raises UnicodeEncodeError, the flush is - # abandoned, and the session silently stops persisting for the - # rest of its life. Scrub so persistence never fails. return _sanitize_surrogates(content) if content is None or isinstance(content, (bytes, int, float)): return content try: - # json.dumps defaults to ensure_ascii=True, which escapes any - # surrogate as \udXXX — already safe to bind. + # ensure_ascii=True escapes surrogates as \\udXXX — safe to bind. return cls._CONTENT_JSON_PREFIX + json.dumps(content) except (TypeError, ValueError): - # Last-resort fallback: stringify so persistence never fails. return _sanitize_surrogates(str(content)) @classmethod def _decode_content(cls, content: Any) -> Any: """Reverse :meth:`_encode_content`; returns scalars unchanged.""" if isinstance(content, str) and content.startswith(cls._CONTENT_JSON_PREFIX): - try: - return json.loads(content[len(cls._CONTENT_JSON_PREFIX):]) - except (json.JSONDecodeError, TypeError): - logger.warning( - "Failed to decode JSON-encoded message content; " - "returning raw string" - ) - return content + return _json_or( + content[len(cls._CONTENT_JSON_PREFIX):], content, + "Failed to decode JSON-encoded message content; returning raw string", + ) return content @staticmethod def _encode_display_metadata(display_metadata: Any) -> Optional[str]: - """Serialize ``display_metadata`` for its TEXT column without double-encoding. - - Import/replace paths can hand us an already-serialized JSON string (the - same hazard ``tool_calls`` guards against above). ``json.dumps`` on that - string would store a quoted JSON string, and the single ``json.loads`` - on read then yields a ``str`` instead of a dict. - """ + """Serialize ``display_metadata`` for its TEXT column without double-encoding an + already-serialized JSON string (import/replace paths hand those in).""" if not display_metadata: return None if isinstance(display_metadata, str): @@ -142,134 +167,15 @@ class SessionMessagesMixin: if isinstance(display_metadata, dict): return json.dumps(display_metadata) logger.warning( - "Ignoring unexpected display metadata type on write: %s", - type(display_metadata).__name__, + "Ignoring unexpected display metadata type on write: %s", type(display_metadata).__name__, ) return None - def _check_transcript_write_guards( - self, - conn, - session_id: str, - compression_lock_holder: Optional[str], - turn_lease_holder: Optional[str] = None, - turn_lease_ttl_seconds: float = 300.0, - 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 :meth:`append_message` and :meth:`append_messages_batch` so - the two writers can never diverge on these correctness invariants - (this guard has already needed targeted fixes — see the #74478 patience - note below). User-initiated transcript mutations may opt in to rejecting - an active unowned turn lease in that same transaction. - """ - from hermes_state import CompressionSessionClosedError, SessionCompressionInProgressError, SessionTurnLeaseLostError, _compression_lock_holder_process_is_dead - # NOTE (#75316 redesign): appends do NOT check compression_locks. - # The lock's job is to stop two COMPRESSIONS colliding, not to fence - # ordinary transcript writes. Concurrent appends during a compression - # are safe by construction: archive_and_compact() commits against a - # watermark captured at compression start and clones every row that - # arrived after it back into the live transcript, in the same write - # transaction. Blocking appends here was the root cause of a whole - # symptom family — turns dying as session_persistence_failed while a - # slow provider summary held the lease (#74568, #77386), including - # stale locks from dead PIDs blocking writes for the full TTL. - # Destructive user mutations are different: a compressor that already - # captured its watermark can otherwise publish the pre-rewind snapshot - # after the mutation and resurrect the removed turn. Keep that narrow - # fence opt-in so ordinary appends retain the watermark behavior. - if reject_active_compression_lock: - active_lock = conn.execute( - "SELECT holder, expires_at FROM compression_locks " - "WHERE session_id = ?", - (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 FROM compression_locks " - "WHERE session_id = ? AND holder = ?", - (session_id, current_holder), - ) - elif current_holder != compression_lock_holder: - raise SessionCompressionInProgressError( - f"Session {session_id!r} is being compressed by another writer" - ) - if turn_lease_holder or reject_active_turn_lease: - conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) - lease = conn.execute( - "SELECT holder, expires_at FROM session_turn_leases " - "WHERE conversation_id = ?", - (conversation_id,), - ).fetchone() - now = time.time() - if turn_lease_holder: - if lease is None or lease["holder"] != turn_lease_holder: - raise SessionTurnLeaseLostError( - f"Session turn lease lost; refusing transcript write " - f"for {session_id!r}" - ) - if float(lease["expires_at"]) <= now: - # Expiry makes the row reclaimable; it does not prove that a - # takeover occurred. BEGIN IMMEDIATE serializes this renewal - # with acquisition, so a still-matching owner can recover from - # a starved refresher without weakening the foreign-holder fence. - conn.execute( - "UPDATE session_turn_leases SET expires_at = ? " - "WHERE conversation_id = ? AND holder = ?", - ( - 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) - ): - # Match acquisition semantics: an expired or provably dead - # owner is reclaimable. Deleting it inside this BEGIN IMMEDIATE - # transaction 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: - raise SessionTurnLeaseLostError( - f"Session has an active turn lease; refusing transcript " - f"mutation for {session_id!r}" - ) - session = conn.execute( - "SELECT ended_at, end_reason FROM sessions WHERE id = ?", - (session_id,), - ).fetchone() - if ( - session is not None - and session["ended_at"] is not None - and session["end_reason"] == "compression" - and not allow_closed_compression_parent - ): - raise CompressionSessionClosedError(session_id) - @staticmethod def _decode_display_metadata(raw: Any) -> Optional[Dict[str, Any]]: - """Decode a ``display_metadata`` column into the dict every reader expects. - - Every message read path must go through this. Returning the raw TEXT - instead reaches the desktop as a string, where ``'task_count' in meta`` - throws and fails the whole resume. Rows written before the encode guard - landed are double-encoded, so unwrap a second layer when we find one. - """ + """Decode a ``display_metadata`` column into a dict (never the raw TEXT — the + desktop does ``'task_count' in meta``). Pre-guard rows are double-encoded, so a + second string layer is unwrapped.""" if raw is None: return None try: @@ -286,25 +192,129 @@ class SessionMessagesMixin: @staticmethod def _reasoning_json_text(value: Any) -> Optional[str]: - """Serialize a structured reasoning field for its TEXT column. - - ``reasoning_details`` / ``codex_reasoning_items`` / ``codex_message_items`` - arrive as list/dict structures from the live runtime, but callers that - round-trip stored rows — ``get_messages`` straight into - ``replace_messages``, e.g. the POST /api/sessions/{id}/fork handler — - hand back the raw TEXT these columns already hold, because - ``get_messages`` only deserializes ``content`` and ``tool_calls``. - Re-dumping that TEXT double-encodes it, and the forked session's next - ``get_messages_as_conversation`` json.loads then yields the inner - string instead of the original list, so every reasoning-replay consumer - (all of which check ``isinstance(..., list)``) silently drops it. - Strings are therefore stored as-is; structures are dumped. - """ + """Serialize a structured reasoning field for its TEXT column. Strings are + stored as-is: round-tripping callers (get_messages -> replace_messages, e.g. the + fork handler) hand back the raw TEXT, and re-dumping would double-encode it so + every reasoning-replay consumer (``isinstance(..., list)``) drops it.""" if not value: return None - if isinstance(value, str): - return value - return json.dumps(value) + return value if isinstance(value, str) else json.dumps(value) + + def _check_transcript_write_guards( + self, + conn, + session_id: str, + compression_lock_holder: Optional[str], + turn_lease_holder: Optional[str] = None, + turn_lease_ttl_seconds: float = 300.0, + 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 + 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: + raise SessionCompressionInProgressError( + f"Session {session_id!r} is being compressed by another writer" + ) + if turn_lease_holder or reject_active_turn_lease: + conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) + lease = conn.execute(_TURN_LEASE_ROW_SQL, (conversation_id,)).fetchone() + now = time.time() + if turn_lease_holder: + if lease is None or lease["holder"] != turn_lease_holder: + raise SessionTurnLeaseLostError( + f"Session turn lease lost; refusing transcript write for {session_id!r}" + ) + if float(lease["expires_at"]) <= now: + # Expiry makes the row reclaimable, it does not prove a takeover; + # BEGIN IMMEDIATE serializes this renewal with acquisition, so a + # still-matching owner recovers from a starved refresher. + conn.execute( + "UPDATE session_turn_leases SET expires_at = ? " + "WHERE conversation_id = ? AND holder = ?", + (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: + raise SessionTurnLeaseLostError( + f"Session has an active turn lease; refusing transcript mutation for {session_id!r}" + ) + 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) + + def _message_row_params( + self, session_id: str, role: str, msg: Dict[str, Any], tool_calls: Any, + message_timestamp: float, *, keep_reasoning: bool, + ) -> tuple: + """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. + """ + 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 + + return ( + session_id, + role, + self._encode_content(msg.get("content")), + msg.get("tool_call_id"), + json.dumps(tool_calls) if tool_calls else None, + _scrub_surrogates(msg.get("tool_name")), + msg.get("effect_disposition"), + message_timestamp, + msg.get("token_count"), + msg.get("finish_reason"), + _scrub_surrogates(_reasoning("reasoning")), + _scrub_surrogates(_reasoning("reasoning_content")), + self._reasoning_json_text(_reasoning("reasoning_details")), + self._reasoning_json_text(_reasoning("codex_reasoning_items")), + self._reasoning_json_text(_reasoning("codex_message_items")), + # `message_id` is yuanbao's existing convention on message dicts. + msg.get("platform_message_id") or msg.get("message_id"), + 1 if msg.get("observed") else 0, + 1 if msg.get("_compressed_summary") else 0, + 1, + _str_or_none(msg.get("api_content")), + _str_or_none(msg.get("display_kind")), + self._encode_display_metadata(msg.get("display_metadata")), + ) def append_message( self, @@ -333,104 +343,40 @@ class SessionMessagesMixin: 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). """ - Append a message to a session. Returns the message row ID. - - Also increments the session's message_count (and tool_call_count - if role is 'tool' or tool_calls is present). - - ``platform_message_id`` is the external messaging platform's own - message ID (e.g. Telegram update_id, Yuanbao msg_id). It is - independent of the SQLite autoincrement primary key and is used by - platform-specific flows like yuanbao's recall guard to redact a - message by its platform-side identifier. - - ``api_content`` is the exact content string sent to the API for this - message when it differs from ``content`` (ephemeral memory/plugin - injections, persist overrides). It is a byte-fidelity sidecar for - prompt-cache-stable replay — stored as sent, except lone surrogates - (which sqlite3 cannot bind and which the conversation loop scrubs - from every outgoing payload anyway, so the scrubbed form IS the - wire bytes). - """ - from hermes_state import _scrub_surrogates - # Display metadata is presentation-only and never changes the model - # context role/content replayed to providers. + 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, + } + # Encode outside the write txn (display metadata first: log-order parity). display_metadata_json = self._encode_display_metadata(display_metadata) - # Serialize structured fields to JSON before entering the write txn - reasoning_details_json = self._reasoning_json_text(reasoning_details) - codex_items_json = self._reasoning_json_text(codex_reasoning_items) - codex_message_items_json = self._reasoning_json_text(codex_message_items) - # tool_calls may arrive as a Python list (from the live agent) or - # as a JSON string (from import/export). Parse first to avoid - # double-encoding. - if isinstance(tool_calls, str): - try: - tool_calls = json.loads(tool_calls) - except (json.JSONDecodeError, TypeError): - tool_calls = [] - tool_calls_json = json.dumps(tool_calls) if tool_calls else None - # Multimodal content (list of parts) must be JSON-encoded: sqlite3 - # cannot bind list/dict parameters directly. - stored_content = self._encode_content(content) - - message_timestamp = time.time() - if timestamp is not None: - try: - if hasattr(timestamp, "timestamp"): - message_timestamp = float(timestamp.timestamp()) - else: - message_timestamp = float(timestamp) - except (TypeError, ValueError): - logger.debug("Ignoring invalid explicit message timestamp: %r", timestamp) - - # Pre-compute tool call count - num_tool_calls = 0 - if tool_calls is not None: - num_tool_calls = len(tool_calls) if isinstance(tool_calls, list) else 1 + msg["display_metadata"] = display_metadata_json + 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, + ) 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, + conn, session_id, compression_lock_holder, + turn_lease_holder=turn_lease_holder, turn_lease_ttl_seconds=turn_lease_ttl_seconds, ) - cursor = conn.execute( - """INSERT INTO messages (session_id, role, content, tool_call_id, - tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason, - reasoning, reasoning_content, reasoning_details, codex_reasoning_items, - codex_message_items, platform_message_id, observed, _compressed_summary, active, api_content, display_kind, display_metadata) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", - ( - session_id, - role, - stored_content, - tool_call_id, - tool_calls_json, - _scrub_surrogates(tool_name), - effect_disposition, - message_timestamp, - token_count, - finish_reason, - _scrub_surrogates(reasoning), - _scrub_surrogates(reasoning_content), - reasoning_details_json, - codex_items_json, - codex_message_items_json, - platform_message_id, - 1 if observed else 0, - 1 if _compressed_summary else 0, - 1, - _scrub_surrogates(api_content) if isinstance(api_content, str) else None, - _scrub_surrogates(display_kind) if isinstance(display_kind, str) else None, - display_metadata_json, - ), - ) - msg_id = cursor.lastrowid - - # Update counters + msg_id = conn.execute(_INSERT_MESSAGE_SQL, params).lastrowid if num_tool_calls > 0: conn.execute( """UPDATE sessions SET message_count = message_count + 1, @@ -439,19 +385,13 @@ class SessionMessagesMixin: ) else: conn.execute( - "UPDATE sessions SET message_count = message_count + 1 WHERE id = ?", - (session_id,), + "UPDATE sessions SET message_count = message_count + 1 WHERE id = ?", (session_id,), ) return msg_id - # Transcript append is THE critical write: its failure aborts the - # user's turn (session_persistence_failed). Use the long patience so - # a sibling process legitimately holding the write lock for seconds - # (VACUUM, TRUNCATE checkpoint at close, an older pre-bounded-merge - # process's FTS optimize) can't destroy a healthy turn (#74478). - return self._execute_write( - _do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S - ) + # THE critical write (its failure aborts the turn): long patience so a sibling + # legitimately holding the lock for seconds (VACUUM, checkpoint) can't kill it. + return self._execute_write(_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S) def append_messages_batch( self, @@ -462,73 +402,40 @@ class SessionMessagesMixin: chunk_rows: Optional[int] = None, turn_lease_ttl_seconds: float = 300.0, ) -> int: - """Append multiple messages atomically in ONE write transaction. + """Append *messages* (``_insert_message_rows`` dict shape) in ONE write txn. - ``messages`` is a list of dicts in the same shape - :meth:`_insert_message_rows` already consumes for replace/compact/ - import (role, content, tool_name, tool_calls, tool_call_id, - finish_reason, reasoning*, codex_*, timestamp, api_content, - display_kind, display_metadata, ...). Reusing that helper keeps ONE - row-serialization path for every multi-row writer. - - A turn-boundary flush writes the whole turn (user + assistant + tool - rows, typically 3-8 messages) as one BEGIN IMMEDIATE / commit pair - instead of one transaction (and, off WAL, one fsync) per row. - - Atomicity contract: all rows land or none do (the caller re-flushes - unstamped messages on the next attempt). The same admission guards - as :meth:`append_message` run once for the batch — same session, - same instant. - - ``chunk_rows`` bounds the transaction size for LARGE copies (branch - seeds can be thousands of rows; measured: 10k rows ≈ 2.4s inside one - BEGIN IMMEDIATE because the FTS triggers run per row, which would - monopolize the write lock and starve concurrent writers). When set, - the batch commits in chunks of at most that many rows — same - recovery semantics as the old per-row loops (a mid-copy failure - leaves a partial seed), just with bounded lock holds. A turn flush - never needs it. Returns the inserted row count. + 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. """ if not messages: return 0 - if chunk_rows is not None and len(messages) > chunk_rows: - inserted_total = 0 - for start in range(0, len(messages), chunk_rows): - inserted_total += self.append_messages_batch( - session_id, - messages[start:start + chunk_rows], + return sum( + self.append_messages_batch( + 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, ) - return inserted_total + 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, + conn, session_id, compression_lock_holder, + 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, + conn, session_id, messages, + encode_content_fn=self._encode_content, decode_content_fn=self._decode_content, ) - inserted = 0 - tool_calls_total = 0 + inserted = tool_calls_total = 0 if inserted_rows: - inserted, tool_calls_total = self._insert_message_rows( - conn, session_id, inserted_rows - ) - - # One aggregated counter update for the newly 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 + ?, @@ -542,22 +449,14 @@ class SessionMessagesMixin: ) return inserted - # Same criticality as append_message: this IS the turn's transcript. - return self._execute_write( - _do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S - ) + return self._execute_write(_do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S) def set_latest_matching_message_display_kind( self, session_id: str, *, role: str, content: str, display_kind: str, display_metadata: Optional[Dict[str, Any]] = None, ) -> bool: - """Stamp presentation metadata on this turn's freshly persisted row. - - The model still receives ``role`` and ``content`` unchanged. Gateway and - CLI synthetic inputs call this immediately after their serial turn has - flushed, preserving producer provenance without classifying by content - during transcript rendering. - """ + """Stamp presentation metadata on this turn's freshly persisted row; the model + still receives ``role``/``content`` unchanged.""" from hermes_state import _scrub_surrogates if not session_id or not content or not display_kind: return False @@ -572,30 +471,25 @@ class SessionMessagesMixin: return False conn.execute( "UPDATE messages SET display_kind = ?, display_metadata = ? WHERE id = ?", - ( - _scrub_surrogates(display_kind), - self._encode_display_metadata(display_metadata), - row[0], - ), + (_scrub_surrogates(display_kind), self._encode_display_metadata(display_metadata), row[0]), ) return True return bool(self._execute_write(_do)) + def _reaction_list(self, meta: Optional[Dict[str, Any]]) -> List[Dict[str, Any]]: + """Well-formed (dict) reactions stored under ``REACTIONS_METADATA_KEY``.""" + reactions = (meta or {}).get(self.REACTIONS_METADATA_KEY) + return [r for r in reactions if isinstance(r, dict)] if isinstance(reactions, list) else [] + def set_message_reaction( - self, - session_id: str, - message_row_id: int, - emoji: Optional[str], - *, - author: str = "user", + 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. - iOS Tapback semantics: one reaction per author per message. Re-sending - the same emoji clears it, a different emoji replaces it. Returns the - message's full reaction list after the write, or ``None`` when the row - doesn't exist or isn't part of *session_id*. + 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*. """ from hermes_state import _scrub_surrogates if not session_id or message_row_id is None: @@ -608,36 +502,17 @@ class SessionMessagesMixin: ).fetchone() if row is None: return None - meta = self._decode_display_metadata(row[0]) or {} - existing = meta.get(self.REACTIONS_METADATA_KEY) - reactions = [ - r - for r in (existing if isinstance(existing, list) else []) - if isinstance(r, dict) and r.get("author") != author - ] - previous = next( - ( - r - for r in (existing if isinstance(existing, list) else []) - if isinstance(r, dict) and r.get("author") == author - ), - None, - ) - # Tapping the live reaction again retracts it. - toggling_off = ( - emoji is not None and previous is not None and previous.get("emoji") == emoji - ) + 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: - reactions.append( - {"emoji": _scrub_surrogates(emoji), "author": author, "at": time.time()} - ) - + reactions.append({"emoji": _scrub_surrogates(emoji), "author": author, "at": time.time()}) if reactions: meta[self.REACTIONS_METADATA_KEY] = reactions else: 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), @@ -646,34 +521,21 @@ class SessionMessagesMixin: return self._execute_write(_do) - def get_message_reactions( - self, session_id: str, message_row_id: int - ) -> List[Dict[str, Any]]: - """Return the reaction list persisted on one message row (never ``None``).""" + def get_message_reactions(self, session_id: str, message_row_id: int) -> List[Dict[str, Any]]: + """Reaction list persisted on one message row (never ``None``).""" if not session_id or message_row_id is None: return [] - row = self._read_one( "SELECT display_metadata FROM messages WHERE id = ? AND session_id = ?", (message_row_id, session_id), ) + return self._reaction_list(self._decode_display_metadata(row[0])) if row is not None else [] - if row is None: - return [] - - meta = self._decode_display_metadata(row[0]) or {} - reactions = meta.get(self.REACTIONS_METADATA_KEY) - - return [r for r in reactions if isinstance(r, dict)] if isinstance(reactions, list) else [] - - def take_unseen_reactions( - self, session_id: str, *, author: str = "user" - ) -> List[Dict[str, Any]]: + 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. - Powers the cache-safe model-context path: reactions are announced on the - NEXT user turn (never by rewriting the message that was reacted to), and - the ``seen`` stamp guarantees each one is announced exactly once. + Reactions are announced on the NEXT user turn (never by rewriting the reacted + message — cache-safe); the ``seen`` stamp makes each announcement exactly once. """ if not session_id: return [] @@ -685,7 +547,6 @@ class SessionMessagesMixin: "ORDER BY id", (session_id,), ).fetchall() - pending = [] for row in rows: meta = self._decode_display_metadata(row["display_metadata"]) @@ -694,33 +555,24 @@ class SessionMessagesMixin: reactions = meta.get(self.REACTIONS_METADATA_KEY) if not isinstance(reactions, list): continue - changed = False for reaction in reactions: - if ( - not isinstance(reaction, dict) - or reaction.get("author") != author - or reaction.get("seen") - ): + if not isinstance(reaction, dict) or reaction.get("author") != author or reaction.get("seen"): continue reaction["seen"] = True 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 "", - } - ) - + pending.append({ + "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"]), ) - return pending return self._execute_write(_do) or [] @@ -730,151 +582,58 @@ class SessionMessagesMixin: ) -> Optional[int]: """Row id of the most recent active message with *role*, or ``None``. - Two callers, same need — "the message I mean, without an id": the agent - defaulting to the turn that triggered it, and the desktop reacting to a - live message that hasn't round-tripped through a resume yet. - ``offset`` steps to earlier turns (1 = the one before the latest) so a - reaction can land retroactively — "two messages ago" is how the caller - thinks about it. - - ``require_text`` (default) skips rows with no plain-text content — - tool-call-only assistant turns and attachment stubs don't render as - bubbles, so "the latest message" as a HUMAN means it must never - resolve to one (a reaction landing on an invisible row looks dropped, - and its annotation quotes an empty string). + ``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. """ 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 "" - ) - + 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)), ) - return row[0] if row else None def latest_user_message_row_id(self, session_id: str) -> Optional[int]: - """Row id of the most recent active user message, or ``None``. - - The agent's default reaction target: "the message that triggered me", - so the model never has to thread row ids through a tool call (mirrors - the photon adapter's ``_record_last_inbound``). - """ + """Row id of the most recent active user message ("the message that triggered + me"), or ``None``.""" return self.latest_message_row_id(session_id, role="user") def get_message_role(self, session_id: str, row_id: int) -> Optional[str]: - """Role of the active message at *row_id* in *session_id*, or ``None``. - - Lets a reaction event carry the target's role so a renderer can match - a live message that doesn't know its durable row id yet. - """ + """Role of the active message at *row_id* in *session_id*, or ``None``.""" if not session_id: return None - row = self._read_one( "SELECT role FROM messages WHERE id = ? AND session_id = ? AND active = 1", (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 for *session_id*. + """Insert *messages* as fresh active rows inside the caller's write txn. - Shared by :meth:`replace_messages` (delete-then-insert) and - :meth:`archive_and_compact` (soft-archive-then-insert). Runs inside the - caller's write transaction (takes the live ``conn``). Returns - ``(inserted_count, tool_call_count)``. Does NOT touch sessions.* counters - — the caller owns that, since the two flows reconcile counts differently. + 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"]``. """ - from hermes_state import _scrub_surrogates now_ts = time.time() - inserted = 0 - tool_calls_total = 0 + inserted = tool_calls_total = 0 for msg in messages: role = msg.get("role", "unknown") - tool_calls = msg.get("tool_calls") - message_timestamp = now_ts - if msg.get("timestamp") is not None: - try: - ts_value = msg.get("timestamp") - if hasattr(ts_value, "timestamp"): - message_timestamp = float(ts_value.timestamp()) - else: - message_timestamp = float(ts_value) - except (TypeError, ValueError): - logger.debug("Ignoring invalid explicit message timestamp: %r", msg.get("timestamp")) - reasoning_details = msg.get("reasoning_details") if role == "assistant" else None - codex_reasoning_items = ( - msg.get("codex_reasoning_items") if role == "assistant" else None - ) - codex_message_items = ( - msg.get("codex_message_items") if role == "assistant" else None - ) - reasoning_details_json = self._reasoning_json_text(reasoning_details) - codex_items_json = self._reasoning_json_text(codex_reasoning_items) - codex_message_items_json = self._reasoning_json_text(codex_message_items) - # tool_calls may arrive as a Python list (from the live agent) - # or as a JSON string (from import_sessions / export_session, - # which store it as TEXT). json.dumps on an already-serialized - # string double-encodes it, so parse first. - if isinstance(tool_calls, str): - try: - tool_calls = json.loads(tool_calls) - except (json.JSONDecodeError, TypeError): - tool_calls = [] - tool_calls_json = json.dumps(tool_calls) if tool_calls else None - # Accept either `platform_message_id` (new explicit name) or - # `message_id` (yuanbao's existing convention on message dicts). - platform_msg_id = ( - msg.get("platform_message_id") or msg.get("message_id") - ) - - api_content = msg.get("api_content") - + tool_calls = _parse_tool_calls(msg.get("tool_calls")) + message_timestamp = _coerce_timestamp(msg.get("timestamp"), now_ts) cur = conn.execute( - """INSERT INTO messages (session_id, role, content, tool_call_id, - tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason, - reasoning, reasoning_content, reasoning_details, codex_reasoning_items, - codex_message_items, platform_message_id, observed, _compressed_summary, active, api_content, display_kind, display_metadata) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", - ( - session_id, - role, - self._encode_content(msg.get("content")), - msg.get("tool_call_id"), - tool_calls_json, - _scrub_surrogates(msg.get("tool_name")), - msg.get("effect_disposition"), - message_timestamp, - msg.get("token_count"), - msg.get("finish_reason"), - _scrub_surrogates(msg.get("reasoning")) if role == "assistant" else None, - _scrub_surrogates(msg.get("reasoning_content")) if role == "assistant" else None, - reasoning_details_json, - codex_items_json, - codex_message_items_json, - platform_msg_id, - 1 if msg.get("observed") else 0, - 1 if msg.get("_compressed_summary") else 0, - 1, - _scrub_surrogates(api_content) if isinstance(api_content, str) else None, - _scrub_surrogates(msg.get("display_kind")) if isinstance(msg.get("display_kind"), str) else None, - self._encode_display_metadata(msg.get("display_metadata")), + _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 - if tool_calls is not None: - tool_calls_total += ( - len(tool_calls) if isinstance(tool_calls, list) else 1 - ) + tool_calls_total += _tool_calls_count(tool_calls) now_ts = max(now_ts + 1e-6, message_timestamp + 1e-6) return inserted, tool_calls_total @@ -886,92 +645,40 @@ class SessionMessagesMixin: archive_dropped: bool = False, reject_active_turn_lease: bool = False, ) -> None: - """Atomically replace the stored messages for a session. + """Atomically replace the stored messages for a session (/retry, /undo, /compress). - Used by transcript-rewrite flows such as /retry, /undo, and /compress. - The delete + reinsert sequence must commit as one transaction so a - mid-rewrite failure does not leave SQLite with a partial transcript. - - DESTRUCTIVE by default: every row for the session is DELETEd (and drops - out of the FTS index). For compaction that must preserve the - pre-compaction transcript under the same id, use - :meth:`archive_and_compact` instead. - - Pass ``active_only=True`` to replace ONLY the live (``active = 1``) rows, - leaving soft-archived rows (``active = 0`` — e.g. the ``compacted = 1`` - turns that :meth:`archive_and_compact` keeps on disk for #38763 - durability, or rewind/undo rows) untouched. Callers that share a session - id with an agent already running in-place compaction must use this so a - full-history rewrite doesn't wipe the rows the agent deliberately - archived. ``message_count``/``tool_call_count`` then track the live set, - matching :meth:`archive_and_compact`. - - Pass ``archive_dropped=True`` to SOFT-archive the live rows instead of - DELETEing them: the replaced turns stay on disk with ``active = 0``, - ``compacted = 0`` — the same "the user took it back" marking - :meth:`rewind_to_message` applies — and stay readable via - :meth:`get_messages` with ``include_inactive=True``. This is the mode a - rewind/edit/regenerate must use: those flows overwrite a transcript the - user may not have meant to drop, and a plain DELETE also evicts the rows - from the FTS index, leaving nothing to recover from (#82756). It implies - active-only handling — already-archived rows are never touched — so - ``active_only`` is redundant with it. The rewritten set is inserted as - fresh active rows exactly as in the destructive path, so the live view - is identical either way; only the durability of the dropped turns - differs. - - Pass ``reject_active_turn_lease=True`` for user-initiated rewrites that - do not already own the cross-process turn lease. The lease check and - transcript mutation then share one write transaction, so a second - process cannot archive or replace a turn that is still being produced. + 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. """ from hermes_state import CompressionSessionClosedError - active_clause = " AND active = 1" if active_only else "" def _do(conn): if reject_active_turn_lease: self._check_transcript_write_guards( - conn, - session_id, - None, - reject_active_turn_lease=True, - reject_active_compression_lock=True, + conn, session_id, None, reject_active_turn_lease=True, reject_active_compression_lock=True, ) - else: - session = conn.execute( - "SELECT ended_at, end_reason FROM sessions WHERE id = ?", - (session_id,), - ).fetchone() - if ( - session is not None - and session["ended_at"] is not None - and session["end_reason"] == "compression" - ): - raise CompressionSessionClosedError(session_id) + elif _ended_by_compression(conn.execute(_ENDED_BY_COMPRESSION_SQL, (session_id,)).fetchone()): + raise CompressionSessionClosedError(session_id) if archive_dropped: - # Content-preserving UPDATE: the rows keep their FTS entries - # (the messages_fts triggers fire on INSERT / DELETE / UPDATE - # of content columns, not on `active`), so the replaced turns - # stay readable via get_messages(include_inactive=True) and - # searchable with include_inactive=True after the rewrite. + # Content-preserving UPDATE: FTS triggers don't fire on `active`, so the + # replaced turns stay searchable/readable with include_inactive=True. conn.execute( - "UPDATE messages SET active = 0 " - "WHERE session_id = ? AND active = 1", - (session_id,), + "UPDATE messages SET active = 0 WHERE session_id = ? AND active = 1", (session_id,), ) else: - conn.execute( - f"DELETE FROM messages WHERE session_id = ?{active_clause}", - (session_id,), - ) + conn.execute(f"DELETE FROM messages WHERE session_id = ?{active_clause}", (session_id,)) conn.execute( - "UPDATE sessions SET message_count = 0, tool_call_count = 0 WHERE id = ?", - (session_id,), - ) - total_messages, total_tool_calls = self._insert_message_rows( - conn, session_id, messages + "UPDATE sessions SET message_count = 0, tool_call_count = 0 WHERE id = ?", (session_id,), ) + 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), @@ -980,35 +687,51 @@ class SessionMessagesMixin: self._execute_write(_do) def has_archived_messages(self, session_id: str) -> bool: - """Return True if the session has any soft-archived (``active = 0``) rows. - - Cheap existence probe — does not load rows. NOTE: production rewrite - paths no longer branch on this (they pass ``active_only=True`` - unconditionally — a probe can fail open or race a concurrent - ``archive_and_compact``, #80216); kept for tests and diagnostics. - """ + """True if the session has any soft-archived (``active = 0``) rows. Cheap probe; + production rewrite paths no longer branch on it (kept for tests/diagnostics).""" return self._read_one( - "SELECT 1 FROM messages WHERE session_id = ? AND active = 0 LIMIT 1", - (session_id,), + "SELECT 1 FROM messages WHERE session_id = ? AND active = 0 LIMIT 1", (session_id,), ) is not None def get_active_message_watermark(self, session_id: str) -> int: - """MAX(id) of the session's active rows — the compression watermark. - - Captured at compression START (before the slow provider summary call). - Every active row with ``id > watermark`` at commit time arrived - concurrently and must survive the compaction verbatim. Returns 0 for - an empty/unknown session. - """ + """MAX(id) of the session's active rows — captured at compression START; every + active row above it arrived concurrently and must survive compaction verbatim. + 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,), + "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]: + """``(ids, tool_call_count)`` of the concurrent-tail rows selected by *sql* + (``SELECT id, tool_calls ...``).""" + rows = conn.execute(sql, params).fetchall() + return [int(r["id"]) for r in rows], sum(_tool_calls_len(r["tool_calls"]) for r in rows) + + def _clone_message_rows(self, conn, tail_ids: List[int], *, session_id: Optional[str] = None) -> None: + """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 ()) + 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], + ) + def archive_and_compact( self, session_id: str, @@ -1018,70 +741,33 @@ class SessionMessagesMixin: lock_holder: Optional[str] = None, tail_count: int = 0, ) -> int: - """Non-destructive in-place compaction for a single durable session id. + """Non-destructive in-place compaction under ONE durable session id. - Soft-archives the active messages (``active = 0``) and inserts - *compacted_messages* as fresh active rows — atomically, in one write - transaction. The conversation keeps ONE session id for life (#38763) - WITHOUT destroying history: + 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. - - The live-context load (:meth:`get_messages_as_conversation`, - :meth:`get_messages`) filters ``active = 1`` by default, so the model - reloads ONLY the compacted set. - - The archived pre-compaction turns stay on disk (active=0) and stay - DISCOVERABLE: they are marked compacted=1, and search_messages() - includes compacted=1 rows by default — so session_search still finds - them, unlike rewind/undo rows (active=0, compacted=0) which stay - hidden. They remain in the FTS index (the messages_fts* triggers - index on INSERT / drop on DELETE and don't key on active/compacted; - flipping to active=0 is a content-preserving UPDATE) and are - recoverable via get_messages(..., include_inactive=True). + *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. - Concurrent-append safety (#75316): when *watermark* is provided (the - value of :meth:`get_active_message_watermark` captured at compression - START), rows that arrived during the slow provider summary call - (``id > watermark``) are NOT summarized away. They are re-sequenced - after the compacted set by a pure-SQL column clone (every column - except ``id`` — content, api_content, platform_message_id, token - counts, reasoning sidecars all survive byte-exact, and the FTS - triggers index the clones naturally), and the originals are archived. - NOTE: re-sequencing assigns the tail rows fresh ids; consumers that - reference durable row ids re-resolve by content (see 3e8ab0610). - ``watermark=None`` preserves the historical archive-everything - behavior. - - Commit-fence safety: when *lock_holder* is provided, the commit - verifies INSIDE the transaction that the compression lock is still - held by that holder and unexpired — a compression whose lease was - reclaimed (crash cleanup, TTL expiry, competing writer) fails the - commit instead of clobbering the winner's transcript. - - *tail_count* (default 0) names how many of the LAST rows of - *compacted_messages* are the verbatim carried-forward tail the - compressor protected rather than summarized (#86366). Those rows' - ORIGINALS — which this call archives as a side effect of the blanket - soft-archive — are superseded byte-identical duplicates, not - "summarized away" content, so they are stamped rewind-style - (``active=0, compacted=0``, hidden from search_messages) instead of - ``compacted=1``. Without this the tail originals satisfy the recall - filter alongside their live clones and session_search returns every - carried-forward message once per compaction. Callers that cannot know - their tail shape keep the historical archive-everything behavior. - - ``message_count`` is set to the ACTIVE count after commit, matching - what the live load returns. ``model_config_patch`` is merged into the - session's JSON config in the same transaction; a ``None`` value - removes that key. Returns the new active count. + ``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. """ from hermes_state import SessionCompressionInProgressError def _do(conn): if lock_holder is not None: - lock_row = conn.execute( - "SELECT holder, expires_at FROM compression_locks " - "WHERE session_id = ?", - (session_id,), - ).fetchone() + lock_row = conn.execute(_COMPRESSION_LOCK_ROW_SQL, (session_id,)).fetchone() if ( lock_row is None or lock_row["holder"] != lock_holder @@ -1091,55 +777,25 @@ class SessionMessagesMixin: f"Compression lease for {session_id!r} lost before " "commit; refusing to publish a stale compaction" ) - patched_model_config = None if model_config_patch is not None: - # on_missing="raise": a prune/compaction must not commit - # against a vanished session row (the compressor's caller - # converts the raised error into a safe keep-the-original - # no-op), unlike the flag setters which tolerate missing rows. + # on_missing="raise": never commit against a vanished session row (the + # compressor's caller turns the error into a keep-the-original no-op). patched_model_config = self._merge_model_config_json( conn, session_id, model_config_patch, on_missing="raise" ) - - # Concurrent tail: active rows that arrived after the watermark. - # Snapshot their ids and tool_calls now — the clone below needs a - # stable id list, and the tool-call count keeps sessions.* honest. tail_ids: list[int] = [] tail_tool_calls = 0 if watermark is not None: - for row in conn.execute( + 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)), - ).fetchall(): - tail_ids.append(int(row["id"])) - raw = row["tool_calls"] - if raw: - try: - parsed = json.loads(raw) if isinstance(raw, str) else raw - tail_tool_calls += len(parsed) if isinstance(parsed, list) else 0 - except (TypeError, ValueError): - pass - - # Soft-archive the live turns: active=0 hides them from the live - # context load, compacted=1 marks them as "summarized away" (vs - # rewind/undo's active=0+compacted=0, which means "user took it - # back"). search_messages includes compacted=1 rows by default so - # the pre-compaction transcript stays discoverable; live-context - # loads (active=1 only) still exclude them. Tail originals whose - # verbatim clones ride inside *compacted_messages* (tail_count) - # are superseded duplicates instead (#86366): they get the - # rewind-style flags so they stop matching the recall filter. - # Rewind-target ids: the originals of the carried-forward tail - # rows (tail_count), captured BEFORE any flag flips. Named apart - # from the watermark `tail_ids` below on purpose — the two are - # different sets (#86366): rewind targets sit AT/BELOW the - # watermark (the compressor only saw rows up to it), while - # `tail_ids` are concurrent appends ABOVE it. Without the bound, - # a concurrent append would steal a LIMIT slot and leave a real - # carried-forward original stamped compacted=1. + ) + # 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 if tail_count > 0: if watermark is not None: @@ -1156,15 +812,9 @@ class SessionMessagesMixin: (session_id, int(tail_count)), ).fetchall() rewind_tail_ids = [int(row["id"]) for row in tail_rows] - - # The watermark clone below re-inserts `tail_ids` rows byte-exact - # as live rows — their originals are the SAME superseded-duplicate - # class as the carried-forward tail (#86366), so they take the - # rewind flags too instead of double-matching the recall filter. rewind_ids = [*(rewind_tail_ids or []), *tail_ids] - if rewind_ids: - placeholders = ",".join("?" for _ in rewind_ids) + placeholders = _placeholders(rewind_ids) conn.execute( "UPDATE messages SET active = 0, compacted = 0 " f"WHERE session_id = ? AND id IN ({placeholders})", @@ -1182,31 +832,11 @@ class SessionMessagesMixin: "WHERE session_id = ? AND active = 1", (session_id,), ) - inserted, tool_calls_total = self._insert_message_rows( - conn, session_id, compacted_messages - ) - + inserted, tool_calls_total = self._insert_message_rows(conn, session_id, compacted_messages) if tail_ids: - # Re-sequence the concurrent tail after the compacted set via - # a pure-SQL column clone: no decode/re-encode round trip, no - # field drift — new id, active=1, compacted=0, all else exact. - placeholders = ",".join("?" for _ in tail_ids) - clone_cols = [ - c for c in self._message_column_names(conn) - if c not in ("id", "active", "compacted") - ] - col_list = ", ".join(clone_cols) - 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, - ) + self._clone_message_rows(conn, tail_ids) inserted += len(tail_ids) tool_calls_total += tail_tool_calls - - # message_count / tool_call_count reflect the LIVE (active) set — - # the archived rows are still on disk but not part of the live count. if model_config_patch is None: conn.execute( "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", @@ -1231,50 +861,33 @@ class SessionMessagesMixin: self._message_columns_cache = cols return cols - def set_latest_user_api_content( - self, session_id: str, content: Any, api_content: str - ) -> int: + 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 (:meth:`archive_and_compact`) inserts the - current turn's user row BEFORE the turn prologue composes the - prefetch/plugin sidecar, and the subsequent crash persist identity-skips - every compacted dict — without this backfill the stamped sidecar would - never land in the DB and any reload would replay clean content, - re-introducing the prompt-cache divergence the sidecar exists to close. - - The ``content`` match is a defensive guard: if the newest active user - row is not the message the caller stamped (racing rewrite, unexpected - tail shape), nothing is written. Returns the number of rows updated - (0 or 1). + 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). """ from hermes_state import _scrub_surrogates - encoded = self._encode_content(content) - 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" ") AND content IS ?", - (_scrub_surrogates(api_content), session_id, encoded), + (_scrub_surrogates(api_content), session_id, self._encode_content(content)), ) def _dedupe_display_generations(self, rows): - """Collapse compaction generations so each message appears once. + """Collapse compaction generations so each logical message appears once. - Compaction epochs copy the protected tail into each new generation, so - one logical message can exist as several rows (identical - role/content/timestamp) with different ``active`` flags and ids. A - display read must surface each exactly once: prefer the live row, then - the newest generation. - - This is the ONE definition shared by every display projection — - :meth:`get_messages` (REST), :meth:`get_resume_conversations` and - :meth:`get_ancestor_display_prefix` (gateway resume), and - :meth:`get_messages_as_conversation` (warm-session payload) — so the - surfaces cannot disagree about the same transcript. *rows* must already - be ordered by ``id``; the returned list keeps that order. + 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. """ seen: Dict[Tuple[Any, ...], Any] = {} for row in rows: @@ -1282,34 +895,50 @@ class SessionMessagesMixin: if row["role"] == "user": from agent.context_compressor import split_user_originated_turn - candidate = { + 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"] - ), - } - handoff, live_view = split_user_originated_turn(candidate) + "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 participate in the dedupe key: compaction copies them - # verbatim, so identical tool messages across generations still - # collapse, while distinct tool calls that happen to share - # role/content/timestamp are never merged. + # 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["role"], dedupe_content, row["timestamp"], + 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 return sorted(seen.values(), key=lambda r: r["id"]) + def _row_to_message_dict(self, row, *, warn_context: str, summary_flag: bool) -> Dict[str, Any]: + """``dict(row)`` with content/tool_calls/display_metadata decoded. *summary_flag* + pops ``_compressed_summary`` and keeps it only as ``True``.""" + msg = dict(row) + if summary_flag and msg.pop("_compressed_summary", 0): + msg["_compressed_summary"] = True + if "content" in msg: + 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 []", + ) + if msg.get("display_metadata") is not None: + msg["display_metadata"] = self._decode_display_metadata(msg["display_metadata"]) + return msg + + @staticmethod + def _active_clause(include_inactive: bool, include_compacted: bool) -> str: + """Audit reads: every row; display reads: active plus compaction-archived + (never Undo/Rewind rows); default: live only.""" + if include_inactive: + return "" + return _DISPLAY_ACTIVE_CLAUSE if include_compacted else " AND active = 1" + def get_messages( self, session_id: str, @@ -1320,69 +949,27 @@ class SessionMessagesMixin: latest: bool = False, after_id: Optional[int] = None, ) -> List[Dict[str, Any]]: - """Load messages for a session in insertion order. + """Load messages for a session in insertion order (AUTOINCREMENT id, never + timestamp — clocks regress on WSL2/NTP steps). - By default only active messages are returned. Pass - ``include_inactive=True`` to load soft-deleted rows (e.g. for - audit / debug views of rewound history). See - :meth:`rewind_to_message` for the soft-delete mechanic. - - Pass ``include_compacted=True`` to additionally load rows preserved - by in-place context compaction (``active=0, compacted=1``). Those are - durable display history, not soft-deleted rows — a user-visible - transcript read must not drop them, or earlier turns silently become - unreachable once the UI exhausts its active-only window. Soft-deleted - Undo/Rewind rows (``active=0, compacted=0``) stay excluded; use - ``include_inactive`` for those. - - Ordered by AUTOINCREMENT id (true insertion order) rather than - timestamp — see c03acca50 for the WSL2 clock-regression rationale. - - When ``limit`` is provided, returns at most ``limit`` messages - starting from ``offset`` (0-based, in insertion order). Enables - pagination for the API endpoint to avoid loading entire transcripts. - With ``latest=True``, the offset is measured back from the newest - message and the selected page is still returned in chronological - order. ``offset`` alone (without ``limit``) also pages — SQLite - requires a LIMIT clause for OFFSET, so it's emitted as ``LIMIT -1`` - (unbounded). - - ``after_id`` enables keyset pagination (``id > after_id``): O(1) - page seeks on huge transcripts where OFFSET degrades to O(n) per - page. Ascending order only (incompatible with ``latest``/``offset``). + ``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. """ 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: raise ValueError("after_id is incompatible with include_compacted (deduped display reads use offset paging)") - if include_inactive: - # Audit / debug reads: every row, including soft-deleted. - active_clause = "" - elif include_compacted: - # Display history: active rows plus rows preserved by in-place - # compaction (active=0, compacted=1), but never soft-deleted - # Undo/Rewind rows (active=0, compacted=0). - active_clause = " AND (active = 1 OR compacted = 1)" - else: - active_clause = " AND active = 1" - keyset_clause = " AND id > ?" if after_id is not None else "" - sql = ( - "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) + active_clause = self._active_clause(include_inactive, include_compacted) if include_compacted: - # Read the full display set (a session's rows are bounded; the - # UI-level 500-row cap lives in the endpoint, not here), dedupe - # generations, then apply paging. - all_rows = self._read_all( - "SELECT * FROM messages WHERE session_id = ?" + active_clause - + " ORDER BY id ASC", + # Read the full display set (the UI-level row cap lives in the endpoint), + # 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], - ) - rows = self._dedupe_display_generations(all_rows) + )) if latest: rows = rows[::-1] rows = rows[offset:] @@ -1391,6 +978,14 @@ class SessionMessagesMixin: if latest: rows = rows[::-1] else: + keyset_clause = " AND id > ?" if after_id is not None else "" + sql = ( + "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) if limit is not None or offset: # SQLite's OFFSET requires LIMIT; -1 means "no limit". sql += " LIMIT ? OFFSET ?" @@ -1398,41 +993,18 @@ class SessionMessagesMixin: rows = self._read_all(sql, params) if latest: rows.reverse() - result = [] - for row in rows: - msg = dict(row) - if msg.pop("_compressed_summary", 0): - msg["_compressed_summary"] = True - if "content" in msg: - msg["content"] = self._decode_content(msg["content"]) - if msg.get("tool_calls"): - try: - msg["tool_calls"] = json.loads(msg["tool_calls"]) - except (json.JSONDecodeError, TypeError): - logger.warning("Failed to deserialize tool_calls in get_messages, falling back to []") - msg["tool_calls"] = [] - if msg.get("display_metadata") is not None: - msg["display_metadata"] = self._decode_display_metadata(msg["display_metadata"]) - result.append(msg) - return result + return [self._row_to_message_dict(row, warn_context="get_messages", summary_flag=True) for row in rows] def find_pr_url_messages(self, session_ids: List[str]) -> List[Dict[str, Any]]: - """Tool results in these sessions that mention a GitHub PR url. - - A candidate scan, deliberately loose: it hands back every tool result - containing ``/pull/`` and leaves the caller to decide which ones make a - claim (see the desktop's PR recovery, which only accepts an output that - is a bare PR url — the signature of ``gh pr create``). Ordered - oldest-first per session so the caller can take the last match. - """ + """Tool results in these sessions containing ``/pull/`` — a deliberately loose + candidate scan, oldest-first per session so the caller can take the last match.""" found: List[Dict[str, Any]] = [] 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] - placeholders = ",".join("?" * len(chunk)) rows = self._read_all( f"""SELECT session_id, content FROM messages - WHERE session_id IN ({placeholders}) + WHERE session_id IN ({",".join("?" * len(chunk))}) AND role = 'tool' AND content LIKE '%/pull/%' ORDER BY id ASC""", chunk, @@ -1440,44 +1012,21 @@ class SessionMessagesMixin: 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]: - """Load a window of messages anchored on a specific message id. + 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). - Returns a dict with: - - ``window``: up to ``window`` messages before the anchor, the anchor - itself, and up to ``window`` messages after, ordered by id ascending. - - ``messages_before``: count of messages strictly before the anchor - still in the session (== window unless we hit the start). - - ``messages_after``: count of messages strictly after the anchor - still in the session (== window unless we hit the end). - - Used by ``session_search`` for both the discovery shape (anchored on the - FTS5 match) and the scroll shape (anchored on any message id). The - ``messages_before`` / ``messages_after`` counts let the caller detect - session boundaries: when either is less than ``window``, the agent has - reached one end of the session. - - Returns an empty window when ``around_message_id`` is not a real id in - ``session_id`` — callers decide how to surface that. + ``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*. """ - if window < 0: - window = 0 + window = max(window, 0) with self._read_ctx() as conn: - # Confirm the anchor exists in this session. 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: return {"window": [], "messages_before": 0, "messages_after": 0} - - # Two queries: anchor + before (DESC, take window+1), and after - # (ASC, take window). Final order is id ASC. before_rows = conn.execute( "SELECT * FROM messages " "WHERE session_id = ? AND id <= ? " @@ -1490,106 +1039,48 @@ class SessionMessagesMixin: "ORDER BY id ASC LIMIT ?", (session_id, around_message_id, window), ).fetchall() - - # before_rows is DESC; reverse so it's ASC, then concatenate after_rows. rows = list(reversed(before_rows)) + list(after_rows) - result = [] - for row in rows: - msg = dict(row) - if "content" in msg: - msg["content"] = self._decode_content(msg["content"]) - if msg.get("tool_calls"): - try: - msg["tool_calls"] = json.loads(msg["tool_calls"]) - except (json.JSONDecodeError, TypeError): - logger.warning( - "Failed to deserialize tool_calls in get_messages_around, falling back to []" - ) - msg["tool_calls"] = [] - if msg.get("display_metadata") is not None: - msg["display_metadata"] = self._decode_display_metadata(msg["display_metadata"]) - result.append(msg) - - # before_rows includes the anchor itself; subtract 1 for the count of - # messages strictly before the anchor in the returned slice. - messages_before = max(0, len(before_rows) - 1) - messages_after = len(after_rows) return { - "window": result, - "messages_before": messages_before, - "messages_after": messages_after, + "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), } def resolve_resume_session_id(self, session_id: str) -> str: """Redirect a resume target to the descendant session that holds the messages. - Context compression ends the current session and forks a new child session - (linked via ``parent_session_id``). The flush cursor is reset, so the - child is where new messages actually land — the parent ends up with - ``message_count = 0`` rows unless messages had already been flushed to - it before compression. See #15000. - - This helper walks ``parent_session_id`` forward from ``session_id`` and - returns the descendant in the chain that has the **most recent** messages. - Unlike the original logic, it does NOT short-circuit when the starting - session already has messages — a descendant that was created by - compression may hold the continuation content and should be preferred - by the WebUI and gateway for ``--resume`` and session loading. - - If no descendant (including the starting session) has any messages, - the original ``session_id`` is returned unchanged. - - The chain is always walked via the child whose ``started_at`` is - latest; that matches the single-chain shape that compression creates. - A depth cap (32) guards against accidental loops in malformed data. + 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. """ if not session_id: return session_id - - # Follow the compression-continuation chain forward to the live tip - # FIRST. Auto-compression ends the current session and forks a - # continuation child, but a long-lived parent keeps its own flushed - # message rows — so the empty-head walk below never redirects it, and - # resuming the parent id reloads the pre-compression transcript while - # the turns generated *after* compression (and their responses) sit in - # the continuation. ``get_compression_tip`` is lineage-aware: it only - # follows children whose parent ended with ``end_reason='compression'`` - # (created after the parent was ended), so delegation / branch children - # never hijack the resume. This is the fix for the desktop "I came back - # and the reply isn't there" report on large sessions. try: tip = self.get_compression_tip(session_id) except Exception: tip = session_id if tip and tip != session_id: session_id = tip - with self._read_ctx() as conn: current = session_id seen = {current} - best = None # tracks the last (deepest) node with messages - + best = None # deepest node with messages for _ in range(32): - # Check if the current node has messages. try: row = conn.execute( - "SELECT 1 FROM messages WHERE session_id = ? LIMIT 1", - (current,), + "SELECT 1 FROM messages WHERE session_id = ? LIMIT 1", (current,), ).fetchone() except Exception: return session_id if row is not None: best = current - - # Walk to the most-recently-started child — but skip explicit - # branch (`_branched_from`), delegate/subagent (`_delegate_from`), - # reset-continuation (`_reset_from` or the legacy same-key - # heuristic — a post-reset conversation must never be reached - # by resuming the parent the user reset away), and tool - # children. They also carry a ``parent_session_id`` yet - # are NOT compression continuations; following them would hijack - # the resume target to an unrelated session (e.g. a subagent - # run). This mirrors the child-exclusion in ``get_compression_tip``. try: child_row = conn.execute( "SELECT id FROM sessions AS child " @@ -1611,9 +1102,20 @@ class SessionMessagesMixin: break seen.add(child_id) current = child_id - return best if best is not None else session_id + def _fetch_conversation_rows(self, session_ids: List[str], active_clause: str, *, with_session_id: bool): + """``_CONVERSATION_ROW_COLUMNS`` rows for *session_ids*, ORDER BY id (insertion + order — timestamps are not monotonic and would break tool-call adjacency).""" + prefix = "SELECT session_id, " if with_session_id else "SELECT " + with self._read_ctx() as conn: + 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), + ).fetchall() + def get_messages_as_conversation( self, session_id: str, @@ -1623,68 +1125,56 @@ class SessionMessagesMixin: include_row_ids: bool = False, include_compacted: bool = False, ) -> List[Dict[str, Any]]: - """ - Load messages in the OpenAI conversation format (role + content dicts). - Used by the gateway to restore conversation history. + """Load messages in OpenAI conversation format (gateway history restore). - By default only active messages are returned. Pass - ``include_inactive=True`` to load soft-deleted (rewound) rows - as well. See :meth:`rewind_to_message`. - - ``include_compacted=True`` additionally loads rows preserved by - in-place compaction (``active=0, compacted=1``), deduped by - :meth:`_dedupe_display_generations`. DISPLAY reads want this; the - model-fed restore must NOT pass it, or a resumed session regrows the - very history compaction just summarized away. - - ``repair_alternation=True`` runs ``repair_message_sequence`` over the - loaded list before returning it. Callers that restore a session for - LIVE REPLAY should pass it: a durable alternation violation (e.g. a - ``user;user`` pair left by a turn that persisted no assistant row) - otherwise re-triggers the pre-request defensive repair on every - single request for the rest of the session's life — the repair - mutates only the per-request list, never the stored transcript. - Inspection/export consumers keep the default and see the transcript - verbatim. + ``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. """ 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) - - if include_inactive: - active_clause = "" - elif include_compacted: - active_clause = " AND (active = 1 OR compacted = 1)" - else: - active_clause = " AND active = 1" - with self._read_ctx() as conn: - placeholders = ",".join("?" for _ in session_ids) - rows = conn.execute( - f"SELECT {self._CONVERSATION_ROW_COLUMNS} " - f"FROM messages WHERE session_id IN ({placeholders})" - # Order by AUTOINCREMENT id (true insertion order), NOT timestamp: - # append_message stamps rows with time.time(), which is not - # monotonic (WSL2, NTP steps, VM/laptop sleep resume). A later - # row can carry an earlier timestamp than its predecessor, and - # ORDER BY timestamp would then sort an assistant tool_calls row - # after its tool response, breaking tool-call/response adjacency - # and triggering an HTTP 400 on replay. This matches get_messages - # — see c03acca50 for the original fix. - f"{active_clause} ORDER BY id", - tuple(session_ids), - ).fetchall() - + rows = self._fetch_conversation_rows( + session_ids, self._active_clause(include_inactive, include_compacted), with_session_id=False, + ) if include_compacted: 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, + rows, session_id=session_id, include_ancestors=include_ancestors, + 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). + """ + 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 + duplicate = None + if previous_exact is not None: + previous_index = next( + (index for index, candidate in enumerate(messages) if candidate is previous_exact), None, + ) + if previous_index is not None: + duplicate = (previous_index, True) + if duplicate is None: + duplicate = self._find_duplicate_replayed_user_message(messages, msg) + if duplicate is not None: + duplicate_index, prefer_current = duplicate + if not prefer_current: + return True, exact_clone_key + messages.pop(duplicate_index) + return False, exact_clone_key + def _rows_to_conversation( self, rows, @@ -1695,175 +1185,72 @@ class SessionMessagesMixin: include_row_ids: bool = False, include_summary_markers: bool = False, ) -> List[Dict[str, Any]]: - """Decode fetched message rows into the OpenAI conversation format. + """Decode fetched message rows (ordered by id, pre-filtered) into OpenAI format. - Extracted from get_messages_as_conversation so get_resume_conversations - can build the model-fed and display views from one SELECT. ``rows`` must - already be ordered by ``id`` (insertion order) and filtered to the - desired session set / active state by the caller. + 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. """ from hermes_state import _strip_background_review_harness, _strip_stale_tool_call_markers messages = [] - # Watermark rotation column-clones concurrent tail rows into the child - # after the new summary, so the copies need not be adjacent. Index the - # exact durable clone identity while decoding instead of rescanning the - # whole accumulated lineage for every user row. exact_user_clones: Dict[Tuple[Any, str], Dict[str, Any]] = {} for row in rows: 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} - # Born durable (#92231): this dict is materialized FROM a durable - # row, so stamp the persistence marker at the source instead of - # relying on every restore caller to thread the loaded list back - # through a flush as ``conversation_history=`` — any - # identity-losing handoff (compression's durable-snapshot - # adoption, incremental persists with no history arg) would - # otherwise re-append the ENTIRE transcript on flush. - # 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 - # Durable per-message identity for surfaces that need to address a - # specific row later (desktop reactions). OPT-IN: only the gateway - # asks for it — every other consumer (ACP restore, export, - # inspection) gets the transcript in its historical shape. - # Underscore-prefixed so every transport's convert_messages() - # strips it before the wire. if include_row_ids and row["id"] is not None: msg["_row_id"] = row["id"] - # api_content is the byte-fidelity sidecar: the exact string sent - # to the API when it differed from the clean content. Returned - # VERBATIM — no sanitize_context, no strip — because the replay - # path substitutes it for content to keep the provider prompt - # cache prefix byte-stable across turns. Cleaning it here would - # re-introduce the divergence it exists to remove. - if row["api_content"]: - msg["api_content"] = row["api_content"] - if row["display_kind"]: - msg["display_kind"] = row["display_kind"] + for col in ("api_content", "display_kind"): + if row[col]: + msg[col] = 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 - if row["timestamp"]: - msg["timestamp"] = row["timestamp"] - if row["tool_call_id"]: - msg["tool_call_id"] = row["tool_call_id"] - if row["tool_name"]: - msg["tool_name"] = row["tool_name"] - if row["effect_disposition"]: - msg["effect_disposition"] = row["effect_disposition"] + for col in ("timestamp", "tool_call_id", "tool_name", "effect_disposition"): + if row[col]: + msg[col] = row[col] if row["tool_calls"]: - try: - msg["tool_calls"] = json.loads(row["tool_calls"]) - except (json.JSONDecodeError, TypeError): - logger.warning("Failed to deserialize tool_calls in conversation replay, falling back to []") - msg["tool_calls"] = [] - # Surface the platform-side message id (e.g. yuanbao msg_id, - # telegram update_id) so platform-specific flows like recall - # can match by external identifier instead of having to fall - # back to content-match heuristics. Exposed as ``message_id`` - # for backward compatibility with the JSONL transcript shape. + msg["tool_calls"] = _json_or( + row["tool_calls"], [], + "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 - # Restore reasoning fields on assistant messages so providers - # that replay reasoning (OpenRouter, OpenAI, Nous) receive - # coherent multi-turn reasoning context. if row["role"] == "assistant": - if row["finish_reason"]: - msg["finish_reason"] = row["finish_reason"] - if row["reasoning"]: - msg["reasoning"] = row["reasoning"] + for col in ("finish_reason", "reasoning"): + if row[col]: + msg[col] = row[col] if row["reasoning_content"] is not None: msg["reasoning_content"] = row["reasoning_content"] - if row["reasoning_details"]: - try: - msg["reasoning_details"] = json.loads(row["reasoning_details"]) - except (json.JSONDecodeError, TypeError): - logger.warning("Failed to deserialize reasoning_details, falling back to None") - msg["reasoning_details"] = None - if row["codex_reasoning_items"]: - try: - msg["codex_reasoning_items"] = json.loads(row["codex_reasoning_items"]) - except (json.JSONDecodeError, TypeError): - logger.warning("Failed to deserialize codex_reasoning_items, falling back to None") - msg["codex_reasoning_items"] = None - if row["codex_message_items"]: - try: - msg["codex_message_items"] = json.loads(row["codex_message_items"]) - except (json.JSONDecodeError, TypeError): - logger.warning("Failed to deserialize codex_message_items, falling back to None") - msg["codex_message_items"] = None + 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", + ) + exact_clone_key = None if include_ancestors: - 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 - ) - duplicate = None - if previous_exact is not None: - previous_index = next( - ( - index - for index, candidate in enumerate(messages) - if candidate is previous_exact - ), - None, - ) - if previous_index is not None: - duplicate = (previous_index, True) - if duplicate is None: - duplicate = self._find_duplicate_replayed_user_message( - messages, msg - ) - if duplicate is not None: - duplicate_index, prefer_current = duplicate - if prefer_current: - # A rotated compression child can carry the same live - # ask as the parent row plus the only surviving summary - # scaffold. Keep the child carrier (and its durable row - # id), not the simpler ancestor copy. - messages.pop(duplicate_index) - else: - continue + skip, exact_clone_key = self._dedupe_replayed_user(messages, msg, exact_user_clones) + if skip: + continue messages.append(msg) if include_ancestors and exact_clone_key is not None: exact_user_clones[exact_clone_key] = msg - # DEFENSE-IN-DEPTH against background-review session pollution: a forked - # skill/memory review that (in older builds, before the _persist_disabled - # fix) shared the parent's session_id wrote its harness turn into this - # real session. The harness is a user/system message instructing the - # agent to "Review the conversation above and update the skill library / - # save to memory" under a hard tool restriction; re-loading it as live - # history makes the agent adopt the curator role and refuse the user's - # actual task. Strip any such harness message AND the curator-mode - # assistant reply immediately following it, so a polluted session - # resumes clean even if stray rows exist. + # Defense-in-depth: strip a background-review harness turn (older builds shared + # the parent's session_id) plus its curator reply, and bare tool-call marker + # content ("[memory]") persisted as an answer before the loop fix. messages = _strip_background_review_harness(messages) - # DEFENSE-IN-DEPTH against #78148: before that fix, a bare tool-call - # marker (e.g. "[memory]") could get cached as a fallback and - # persisted as if it were the model's real answer. Sessions written - # before the fix can still carry those rows — clear the stray - # content on load so replaying history doesn't re-teach the model - # to keep emitting the marker. No-op for unaffected sessions. messages = _strip_stale_tool_call_markers(messages) if repair_alternation and messages: - # Lazy import: hermes_state already depends on agent.* (see - # sanitize_context above), but keep this optional path from - # widening the import surface at module load. from agent.agent_runtime_helpers import repair_message_sequence repaired = repair_message_sequence(None, messages) @@ -1877,145 +1264,62 @@ class SessionMessagesMixin: ) return messages - def get_resume_conversations( - self, session_id: str - ) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]: - """Return ``(model_history, display_history)`` for a session resume in ONE SELECT. + 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. - ``session.resume`` needs two projections of the same lineage: - - - ``model_history`` — the tip session's active rows, alternation-repaired - (the live-replay working conversation). Equivalent to - ``get_messages_as_conversation(session_id, repair_alternation=True)``. - - ``display_history`` — the full compression lineage (ancestors → tip), - verbatim, with replayed-user dedup. Explicit ``/branch`` sessions are - excluded from this lineage because their own rows already contain the - copied transcript; including the live parent's rows would let messages - written to the original after the fork leak into the branch. - - The display projection also includes rows preserved by IN-PLACE - compaction (``active=0, compacted=1``), deduped by - :meth:`_dedupe_display_generations`. Without them a compacted - conversation resumes showing only its summary plus the carried-forward - tail — the user's own turns read as deleted even though every row is - still on disk, and the REST transcript read (which has always included - them) disagreed with this one about the same session (#92080). - - The display fetch already reads a superset of the model fetch (the tip - rows are part of the lineage), so serving both from one lineage SELECT - halves the resume's DB work versus two separate calls, with byte-identical - output (see test_get_resume_conversations_matches_separate_reads). + ``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). """ session_ids = self._resume_lineage_ids(session_id) - with self._read_ctx() as conn: - placeholders = ",".join("?" for _ in session_ids) - rows = conn.execute( - f"SELECT session_id, {self._CONVERSATION_ROW_COLUMNS} " - f"FROM messages WHERE session_id IN ({placeholders}) " - # Compaction-archived rows (active=0, compacted=1) are display - # history; Undo/Rewind rows (active=0, compacted=0) are not. - "AND (active = 1 OR compacted = 1) " - # ORDER BY id (insertion order) — see get_messages_as_conversation - # for why timestamp ordering is unsafe. - "ORDER BY id", - tuple(session_ids), - ).fetchall() - - # Tip rows are exactly the model-fed set (get_messages_as_conversation - # with session_ids=[session_id]); filtering the lineage fetch preserves - # their relative id order. The model projection stays active-only — it - # is the compressed working context and must not regrow the history - # compaction just summarized away. + 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, - # Pre-compress checkpointing: the resumed model history must keep - # the summary marker so checkpoint providers can exclude derivative - # summaries after a process restart (marker survives restart). - include_summary_markers=True, + tip_rows, session_id=session_id, include_ancestors=False, repair_alternation=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, + self._dedupe_display_generations(rows), session_id=session_id, + 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 for *session_id*. - - Compression continuations need their ended ancestors' rows for the - display transcript; an explicit ``/branch`` copy already owns its - transcript, so its lineage is itself alone. This is the ONE definition - shared by the resume readers (``get_resume_conversations``, - ``get_ancestor_display_prefix``) and the resume guard - (``assert_resume_safe`` / ``get_resume_message_count``) — the guard must - count exactly the rows a resume would load, never a superset. - """ + """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) - def get_resume_message_count( - self, session_id: str, *, tip_only: bool = False - ) -> int: - """Count the rows that a resume would materialize. + 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 + DISPLAY set (active + compaction-archived) that get_resume_conversations loads.""" + if tip_only: + return [session_id], "active = 1" + return self._resume_lineage_ids(session_id), "(active = 1 OR compacted = 1)" - ``tip_only=True`` counts the tip segment's ACTIVE rows — the set a - model-history restore loads (``get_messages_as_conversation`` without - ancestors, or the deferred Desktop resume that pages the display - transcript over REST and never materializes the ancestor prefix in - memory). - - Otherwise this counts the full-lineage DISPLAY set — active rows plus - the compaction-archived rows ``get_resume_conversations`` now loads - for the transcript. Counting only active rows here would let a - heavily-compacted conversation pass a limit sized for a handful of - live rows and then materialize tens of thousands. - """ - session_ids = [session_id] if tip_only else self._resume_lineage_ids(session_id) - active_clause = "active = 1" if tip_only else "(active = 1 OR compacted = 1)" - placeholders = ",".join("?" for _ in session_ids) + def get_resume_message_count(self, session_id: str, *, tip_only: bool = False) -> int: + """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}) AND {active_clause}", + f"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 resume row count or reject a transcript too large to load. + 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`` resolves the limit from config - (``sessions.max_resume_messages``); 0 disables the guard and returns - the (bounded) count without raising. - - ``tip_only=True`` bounds only the tip segment's ACTIVE rows, for - callers that never materialize the ancestor lineage or the - compaction archive in memory (tip-only model restore, deferred - Desktop resume whose display history is REST-paginated). A - heavily-compressed conversation — 85 compaction segments and ~29k - lineage rows behind a ~700-row tip — is exactly the shape compression - is supposed to produce; counting its whole lineage against a limit - sized for in-memory materialization rejected the healthiest sessions - (Desktop Bot Chat stuck on "Waking up…" with code 4130) while the - process would only ever have held the tip. - - The full (non-``tip_only``) bound counts the DISPLAY set — active plus - compaction-archived rows — because that is what - ``get_resume_conversations`` materializes for the transcript. + ``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. """ from hermes_state import SessionResumeTooLargeError, resolved_max_resume_messages if max_messages is None: @@ -2023,17 +1327,11 @@ class SessionMessagesMixin: if max_messages < 0: raise ValueError("max_messages must be non-negative") if max_messages == 0: - # Guard disabled by config — skip counting entirely. Every live - # caller invokes this for its raise side effect and ignores the - # return value, and an unbounded lineage COUNT here would do the - # exact pathological work the disable exists to avoid. return 0 - session_ids = [session_id] if tip_only else self._resume_lineage_ids(session_id) - active_clause = "active = 1" if tip_only else "(active = 1 OR compacted = 1)" - placeholders = ",".join("?" for _ in session_ids) + session_ids, active_clause = self._resume_count_scope(session_id, tip_only) row = self._read_one( "SELECT COUNT(*) FROM (" - f"SELECT 1 FROM messages WHERE session_id IN ({placeholders}) " + f"SELECT 1 FROM messages WHERE session_id IN ({_placeholders(session_ids)}) " f"AND {active_clause} LIMIT ?" ")", (*session_ids, max_messages + 1), @@ -2041,117 +1339,66 @@ class SessionMessagesMixin: message_count = int(row[0] if row else 0) if message_count > max_messages: raise SessionResumeTooLargeError( - message_count, - max_messages, + message_count, max_messages, 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]]: - """Return the ancestor-only display messages for a session lineage. + """Ancestor-only display messages of a lineage (rows with ``session_id !=`` tip). - These are messages from parent/grandparent sessions (compression - ancestors) that appear in the display transcript but NOT in the - tip session's model-fed history. Used by ``session.resume`` to - build the ``display_history_prefix`` that ``_live_session_payload`` - prepends to the live model history. - - Previously the prefix was calculated as - ``display_history[:len(display) - len(raw)]``, but that overcounts - when ``repair_message_sequence`` removes messages from the MIDDLE - of the tip history (e.g. verification candidates collapsed by the - consecutive-assistant merge) — the length difference includes both - ancestor messages AND repair-removed tip messages, but the slice - only captures the first N display messages (which are tip messages - when there are no ancestors), causing duplication. This method - returns ONLY the genuine ancestor messages, identified by - ``session_id != tip_session_id``. (#65919) + ``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. """ session_ids = self._resume_lineage_ids(session_id) if len(session_ids) <= 1: return [] - with self._read_ctx() as conn: - placeholders = ",".join("?" for _ in session_ids) - rows = conn.execute( - f"SELECT session_id, {self._CONVERSATION_ROW_COLUMNS} " - f"FROM messages WHERE session_id IN ({placeholders}) " - # Display read: compaction-archived rows included, Undo/Rewind - # rows excluded (see get_resume_conversations). - "AND (active = 1 OR compacted = 1) " - "ORDER BY id", - tuple(session_ids), - ).fetchall() - rows = self._dedupe_display_generations(rows) - ancestor_ids = { - int(row["id"]) - for row in rows - if row["session_id"] != session_id and row["id"] is not None - } + rows = self._dedupe_display_generations( + self._fetch_conversation_rows(session_ids, _DISPLAY_ACTIVE_CLAUSE, with_session_id=True) + ) + ancestor_ids = {int(row["id"]) for row in rows if row["session_id"] != session_id and row["id"] is not None} if not ancestor_ids: return [] lineage = self._rows_to_conversation( - rows, - session_id=session_id, - include_ancestors=True, - repair_alternation=False, - include_row_ids=True, + 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") not in ancestor_ids: - continue - projected = message.copy() - projected.pop("_row_id", None) - prefix.append(projected) + if message.get("_row_id") in ancestor_ids: + projected = message.copy() + projected.pop("_row_id", None) + prefix.append(projected) return prefix def get_conversation_root(self, session_id: str) -> str: - """Return the ROOT id of *session_id*'s lineage chain. - - The root is the stable "conversation id": context compression - rotates ``session_id`` to a new segment linked via - ``parent_session_id``, and delegate subagents hang off their - parent the same way. Walking to the root gives every segment of - one user-facing conversation (and its delegation tree) a single - identifier — used for Nous Portal ``conversation=`` usage tagging. - Returns *session_id* unchanged when it has no recorded parent. - """ + """ROOT id of *session_id*'s lineage — the stable conversation id across + compression segments and delegate subagents (Nous Portal usage tagging). + Unchanged when there is no recorded parent.""" chain = self._session_lineage_root_to_tip(session_id) - return (chain[0] if chain and chain[0] else session_id) + return chain[0] if chain and chain[0] else session_id @staticmethod - def _canonical_replayed_user_content( - msg: Dict[str, Any], - ) -> Tuple[Any, bool]: + def _canonical_replayed_user_content(msg: Dict[str, Any]) -> Tuple[Any, bool]: """Return canonical live content and whether *msg* is composite.""" 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"), + live_view.get("content") if is_composite and live_view is not None else msg.get("content"), is_composite, ) @staticmethod - def _exact_replayed_user_clone_key( - timestamp: Any, content: Any - ) -> Optional[Tuple[Any, str]]: + def _exact_replayed_user_clone_key(timestamp: Any, content: Any) -> Optional[Tuple[Any, str]]: """Return a hashable key for a column-exact rotation clone.""" if timestamp is None or content in (None, "", []): return None try: - encoded = json.dumps( - content, - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - ) + encoded = json.dumps(content, ensure_ascii=False, sort_keys=True, separators=(",", ":")) except (TypeError, ValueError): return None return timestamp, encoded @@ -2162,48 +1409,32 @@ class SessionMessagesMixin: ) -> Optional[Tuple[int, bool]]: """Return an adjacent replay duplicate and whether *msg* must win. - Compression rotation may persist the current ask once in the parent - and again inside a composite child carrier. Compare the canonical live - payload for that carrier, while retaining the historical exact-string - dedupe for ordinary replayed users. The child carrier wins because it - owns both the current durable row identity and the retained scaffold. + 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). """ from hermes_state import SessionDB if msg.get("role") != "user": return None - content, prefer_current = SessionDB._canonical_replayed_user_content(msg) if content in (None, "", []): return None - for index in range(len(messages) - 1, -1, -1): prev = messages[index] if prev.get("role") == "user": - prev_content, prev_is_composite = ( - SessionDB._canonical_replayed_user_content(prev) - ) - if prev_content == content and ( - prefer_current - or prev_is_composite - or isinstance(content, str) - ): + prev_content, prev_is_composite = SessionDB._canonical_replayed_user_content(prev) + if prev_content == content and (prefer_current or prev_is_composite or isinstance(content, str)): return index, prefer_current if prev.get("role") == "assistant" and (prev.get("content") or prev.get("tool_calls")): return None return None def get_active_message_ids(self, session_id: str) -> List[int]: - """Return the ordered physical ids pinned by rewind CAS checks. - - Conversation projections intentionally omit legacy background-review - harness rows. Destructive rewinds must nevertheless pin every active - physical row so the caller snapshot matches the transaction-local - comparison in :meth:`rewind_to_message`. - """ + """Ordered physical active ids pinned by rewind CAS checks (includes legacy + harness rows that conversation projections omit).""" rows = self._read_all( - "SELECT id FROM messages " - "WHERE session_id = ? AND active = 1 ORDER BY id", - (session_id,), + "SELECT id FROM messages WHERE session_id = ? AND active = 1 ORDER BY id", (session_id,), ) return [int(row[0]) for row in rows] @@ -2211,9 +1442,7 @@ class SessionMessagesMixin: def _active_transcript_counts(conn, session_id: str) -> tuple[int, int]: """Return active message/tool-call counts inside the caller's txn.""" rows = conn.execute( - "SELECT tool_calls FROM messages " - "WHERE session_id = ? AND active = 1", - (session_id,), + "SELECT tool_calls FROM messages WHERE session_id = ? AND active = 1", (session_id,), ).fetchall() tool_call_count = 0 for row in rows: @@ -2230,6 +1459,32 @@ class SessionMessagesMixin: tool_call_count += 1 return len(rows), tool_call_count + 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*. + """ + 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")) + handoff, live_view = split_user_originated_turn(split_target) + if live_view is None: + raise ValueError("rewind target is not a user-originated turn") + live_content = live_view.get("content") + if isinstance(live_content, str): + live_content = sanitize_context(live_content).strip() + if expected_target_content is not None and live_content != expected_target_content: + raise RuntimeError("rewind target changed before it could be persisted") + if preserve_compaction_handoff and handoff is None: + raise ValueError("preserve_compaction_handoff requires an active composite carrier") + return handoff if preserve_compaction_handoff else None + def rewind_to_message( self, session_id: str, @@ -2239,167 +1494,75 @@ class SessionMessagesMixin: expected_active_ids: Optional[List[int]] = None, expected_target_content: Any = None, ) -> Dict[str, Any]: - """Soft-delete all messages with id >= ``target_message_id`` in *session_id*. + """Soft-delete (``active=0``) every message with id >= *target_message_id*. - The target message itself becomes inactive as well so the caller - can pre-fill it as the next user prompt without it appearing - twice in the replayed transcript. Rewound rows are kept on - disk with ``active=0`` for audit / forensic inspection — use - :meth:`get_messages` with ``include_inactive=True`` to see them. + 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. - Returns a dict:: - - { - "rewound_count": int, # number of rows newly flipped to active=0 - "target_message": dict, # full row dict of the target - "new_head_id": int|None # id of the last still-active row, or None - } - - Raises ``ValueError`` if the target message does not exist in - *session_id* or if its role is not ``"user"``. With - ``preserve_compaction_handoff=True``, a composite summary carrier is - split inside the same write transaction: its original row is archived - and its canonical hidden handoff scaffold is inserted as the new head. - That opt-in result also contains ``replacement_message_id``. - - ``expected_active_ids`` optionally pins the ordered active row set. - ``expected_target_content`` additionally pins the selected canonical - live-user payload. Both checks run inside the write transaction before - any row or counter mutation. Presentation-only metadata changes (for - example Desktop reactions) deliberately do not invalidate a rewind. - A live cross-process turn lease always refuses the rewind; expired or - provably dead holders are reclaimed inside the mutation transaction. - - Always increments ``sessions.rewind_count`` — even when the - target is already inactive — so the counter accurately reflects - the number of rewind operations performed against the session. - Idempotent on the ``active`` flag: re-rewinding past the same - target is a no-op on row state but still bumps the counter. + ``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. """ def _do(conn): - # Rewind changes the active transcript and must honor the same - # compression/closed-parent and cross-process turn guards as - # append writers. self._check_transcript_write_guards( - conn, - session_id, - None, - reject_active_turn_lease=True, - reject_active_compression_lock=True, + conn, session_id, None, reject_active_turn_lease=True, reject_active_compression_lock=True, ) - if expected_active_ids is not None: active_rows = conn.execute( - "SELECT id FROM messages " - "WHERE session_id = ? AND active = 1 ORDER BY id", - (session_id,), + "SELECT id FROM messages WHERE session_id = ? AND active = 1 ORDER BY id", (session_id,), ).fetchall() - active_ids = [int(active_row[0]) for active_row in active_rows] - if active_ids != expected_active_ids: - raise RuntimeError( - "active transcript changed before the rewind could be persisted" - ) - + if [int(active_row[0]) for active_row 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), + "SELECT * FROM messages WHERE id = ? AND session_id = ?", (target_message_id, session_id), ).fetchone() if row is None: - raise ValueError( - f"message {target_message_id} not found in session {session_id}" - ) + raise ValueError(f"message {target_message_id} not found in session {session_id}") target_row = dict(row) if target_row.get("role") != "user": raise ValueError( 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 if preserve_compaction_handoff or expected_target_content is not None: - 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") - ) - handoff, live_view = split_user_originated_turn(split_target) - if live_view is None: - raise ValueError("rewind target is not a user-originated turn") - live_content = live_view.get("content") - if isinstance(live_content, str): - live_content = sanitize_context(live_content).strip() - if ( - expected_target_content is not None - and live_content != expected_target_content - ): - raise RuntimeError( - "rewind target changed before it could be persisted" - ) - if preserve_compaction_handoff and handoff is None: - raise ValueError( - "preserve_compaction_handoff requires an active composite carrier" - ) - replacement = handoff if preserve_compaction_handoff else None - + replacement = self._split_rewind_target(target_row, expected_target_content, preserve_compaction_handoff) cursor = conn.execute( - "SELECT id FROM messages " - "WHERE session_id = ? AND id >= ? AND active = 1", + "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()] if ids: - placeholders = ",".join("?" for _ in ids) - conn.execute( - f"UPDATE messages SET active = 0 WHERE id IN ({placeholders})", - ids, - ) + conn.execute(f"UPDATE messages SET active = 0 WHERE id IN ({_placeholders(ids)})", ids) if replacement is not None: self._insert_message_rows(conn, session_id, [replacement]) - inserted = conn.execute("SELECT last_insert_rowid()").fetchone() - replacement_message_id = int(inserted[0]) + replacement_message_id = int(conn.execute("SELECT last_insert_rowid()").fetchone()[0]) conn.execute( - "UPDATE sessions SET rewind_count = COALESCE(rewind_count, 0) + 1 " - "WHERE id = ?", - (session_id,), - ) - message_count, tool_call_count = self._active_transcript_counts( - conn, session_id + "UPDATE sessions SET rewind_count = COALESCE(rewind_count, 0) + 1 WHERE id = ?", (session_id,), ) + message_count, tool_call_count = self._active_transcript_counts(conn, session_id) conn.execute( - "UPDATE sessions SET message_count = ?, tool_call_count = ? " - "WHERE id = ?", + "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE 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,), + "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 - ) + 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 - target_row, rewound, new_head_id, replacement_message_id = ( - self._execute_write(_do) - ) - - # Decode content for callers (prefill the prompt buffer) without a - # second fallible database operation after the transaction commits. + 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. target_row["content"] = self._decode_content(target_row.get("content")) - - result = { - "rewound_count": len(rewound), - "target_message": target_row, - "new_head_id": new_head_id, - } + result = {"rewound_count": len(rewound), "target_message": target_row, "new_head_id": new_head_id} if preserve_compaction_handoff: result["replacement_message_id"] = replacement_message_id return result @@ -2408,40 +1571,26 @@ class SessionMessagesMixin: """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,) - ) + cursor = conn.execute("SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,)) else: cursor = conn.execute("SELECT COUNT(*) FROM messages") return cursor.fetchone()[0] - def has_platform_message_id( - self, session_id: str, platform_message_id: str - ) -> bool: - """Check if a message with the given platform_message_id exists. - - Uses the idx_messages_platform_msg_id partial index for efficient - lookup. Used by the gateway's transient-failure dedupe guard (#47237) - to skip re-persisting a user message that was already saved on a - prior retry of the same inbound platform message. - """ + def has_platform_message_id(self, session_id: str, platform_message_id: str) -> bool: + """True when a message with *platform_message_id* exists (partial index lookup; + the gateway's transient-failure dedupe guard).""" return self._read_one( - "SELECT 1 FROM messages " - "WHERE session_id = ? AND platform_message_id = ? LIMIT 1", + "SELECT 1 FROM messages WHERE session_id = ? AND platform_message_id = ? LIMIT 1", (session_id, platform_message_id), ) 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. + """True when *session* is a branch, delegate, or tool child of its parent. - Markers only count as a fork when they point at ``parent_session_id``. - Compression copies ``model_config`` onto the continuation - (``publish_compression_child`` callers pass - ``agent._session_init_model_config``), so a delegate's continuation - carries ``_delegate_from=``. Presence-only - matching would treat that real continuation as a fork — the same - misclassification ``_NON_CONTINUATION_CHILD_FILTER_SQL`` already - avoids by binding both markers to the queried 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``). """ if session.get("source") == "tool": return True @@ -2462,63 +1611,26 @@ class SessionMessagesMixin: return branched is not None or delegated is not None def is_explicit_fork_child(self, session_id: str) -> bool: - """True when ``session_id`` is a /branch, delegate, or tool child row. - - Read-only public view of :meth:`_is_explicit_fork_child_row` for - callers that must respect the fork boundary without re-implementing - its marker rules (``agent/prompt_cache_scope.py`` keeps a declared - conversation key from crossing it). A missing row is not a fork. - """ + """Public read-only view of :meth:`_is_explicit_fork_child_row`; a missing row + is not a fork.""" session = self.get_session(session_id) return bool(session and self._is_explicit_fork_child_row(session)) - def latest_conversation_boundary( - self, session_key: str, source: str - ) -> Optional[int]: - """How many conversation boundaries this routing peer has crossed. + def latest_conversation_boundary(self, session_key: str, source: str) -> Optional[int]: + """How many conversation boundaries (``_RESET_END_REASONS`` ends) this routing + peer has crossed, or ``None`` when never reset. - A boundary is a row this peer ended at an intentional conversation - break — the ``_RESET_END_REASONS`` set (``/new``, ``/switch``, idle, - daily, suspended, resume_pending_expired). That is the same fence - :meth:`find_latest_gateway_session_for_peer` refuses to reach behind, - so the two agree on where one conversation stops and the next begins - and cannot drift. - - The peer is ``(session_key, source)``, the SAME identity tuple recovery - uses — never the key alone. ``X-Hermes-Session-Key`` accepts any - authenticated caller-supplied string, so an API conversation may - legally carry the same key as a Telegram row in one database; keying - on the string alone would let a ``/new`` on that unrelated row rotate - this conversation's affinity identity while recovery correctly refuses - to cross the same line. - - Returns the count, or ``None`` when this peer has never been reset. - - The value comes from ``conversation_generations``, which - :meth:`_bump_conversation_generation` advances inside the transaction - that writes each boundary — NOT from an aggregate over the session - rows. An aggregate cannot prove non-reuse: ``delete_session()`` - orphans children and deletes the row, and bulk prune selects ended - rows, so ``COUNT``/``MAX`` over boundaries can return a pair it already - emitted and hand a new conversation a retired affinity identity. It is - also wall-clock-free, so a backwards NTP correction cannot reorder it. - - Databases upgraded mid-conversation start at no generation and take - their first one from the next boundary written; a conversation that - reset before the upgrade shares its predecessor's scope once, which - costs a warm prompt-cache bucket and never crosses an identity. - - These rows are never garbage-collected, by design: dropping one resets - the peer to "no generation", so its next boundary writes ``1`` again - and re-issues a scope a retired conversation already used — the ABA - this counter exists to prevent. See the schema comment in - ``hermes_state_common.py``. + The peer is ``(session_key, source)`` — the identity recovery uses — never the + key alone (an API caller may legally reuse a Telegram row's key). Read from + ``conversation_generations`` (advanced inside each boundary's txn), not an + aggregate over session rows: deletes/prunes would let an aggregate re-emit a + retired pair. Rows are never garbage-collected, by design (dropping one would + re-issue generation 1 — the ABA this counter prevents). """ if not session_key or not source: return None row = self._read_one( - "SELECT generation FROM conversation_generations " - "WHERE source = ? AND session_key = ?", + "SELECT generation FROM conversation_generations WHERE source = ? AND session_key = ?", (source, session_key), ) if row is None or row["generation"] is None: @@ -2529,48 +1641,21 @@ class SessionMessagesMixin: def clear_messages(self, session_id: str) -> None: """Delete all messages for a session and reset its counters.""" def _do(conn): + conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,)) conn.execute( - "DELETE FROM messages WHERE session_id = ?", (session_id,) - ) - conn.execute( - "UPDATE sessions SET message_count = 0, tool_call_count = 0 WHERE id = ?", - (session_id,), + "UPDATE sessions SET message_count = 0, tool_call_count = 0 WHERE id = ?", (session_id,), ) 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 in the ``messages`` table by sessions persisted before the - #78148 fix in ``agent.conversation_loop``. + 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. - ``_strip_stale_tool_call_markers`` already repairs this in memory on - every session load (see ``_rows_to_conversation``), so running this - is optional — but for long-lived sessions the same rows get - re-scanned and re-repaired on every resume, which is wasted work - and keeps the contaminated bytes sitting in the DB (and in any - downstream cache/backup snapshot of it) indefinitely. This rewrites - the affected rows once, in place. - - Only the ``content`` column is touched — ``role``, ``tool_calls``, - and every other column on the row are left exactly as they are, so - provider tool_call/tool_result pairing is unaffected. - - Unlike the in-memory repair, this UPDATE is permanent and can't be - undone from within the DB. Since ``backup`` defaults to True, a - timestamped full snapshot is taken via ``VACUUM INTO`` (safe against - a live connection, unlike the raw-copy ``_backup_db_file`` used for - malformed-schema repair) before any row is touched — mirroring - ``repair_state_db_schema``'s backup-by-default convention for - destructive state.db operations. No snapshot is taken when there is - nothing to change. - - With ``dry_run=True``, reports the affected row count/ids without - writing or backing up (read-only, no write lock taken). - - Returns ``{"dry_run": bool, "rows_affected": int, "row_ids": [...], - "backup_path": str|None}``. + 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"}``. """ from hermes_state import _STALE_TOOL_CALL_MARKER_RE @@ -2579,40 +1664,24 @@ class SessionMessagesMixin: "SELECT id, content FROM messages " "WHERE role = 'assistant' AND tool_calls IS NOT NULL AND tool_calls != ''" ) - affected: List[int] = [] - for row in cursor.fetchall(): - content = row["content"] - if isinstance(content, str) and _STALE_TOOL_CALL_MARKER_RE.fullmatch(content.strip()): - affected.append(row["id"]) - return affected + return [ + row["id"] for row in cursor.fetchall() + if isinstance(row["content"], str) and _STALE_TOOL_CALL_MARKER_RE.fullmatch(row["content"].strip()) + ] + + def _result(affected, backup_path=None): + return {"dry_run": dry_run, "rows_affected": len(affected), "row_ids": affected, "backup_path": backup_path} with self._read_ctx() as conn: affected_ids = _find_affected(conn) - - if dry_run: - return { - "dry_run": True, - "rows_affected": len(affected_ids), - "row_ids": affected_ids, - "backup_path": None, - } - - if not affected_ids: - return { - "dry_run": False, - "rows_affected": 0, - "row_ids": [], - "backup_path": None, - } - + if dry_run or not affected_ids: + return _result(affected_ids) 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}" - ) + dest = self.db_path.with_name(f"{self.db_path.name}.pre-clean-markers-backup-{stamp}") with self._lock: self._conn.execute("VACUUM INTO ?", (str(dest),)) backup_path = str(dest) @@ -2621,22 +1690,12 @@ class SessionMessagesMixin: def _do(conn): ids = _find_affected(conn) if ids: - placeholders = ",".join("?" * len(ids)) - conn.execute( - f"UPDATE messages SET content = '' WHERE id IN ({placeholders})", - ids, - ) + conn.execute(f"UPDATE messages SET content = '' WHERE id IN ({','.join('?' * len(ids))})", ids) return ids affected_ids = self._execute_write(_do) if affected_ids: logger.info( - "Permanently cleared %d stale tool-call marker row(s) in state.db (#78148)", - len(affected_ids), + "Permanently cleared %d stale tool-call marker row(s) in state.db (#78148)", len(affected_ids), ) - return { - "dry_run": False, - "rows_affected": len(affected_ids), - "row_ids": affected_ids, - "backup_path": backup_path, - } + return _result(affected_ids, backup_path) diff --git a/hermes_state_titles.py b/hermes_state_titles.py index 3d6dbea0f5..9404fa74ad 100644 --- a/hermes_state_titles.py +++ b/hermes_state_titles.py @@ -13,66 +13,44 @@ from hermes_state_common import _COMPRESSION_CHILD_SQL, escape_like as _escape_l # caplog tests pin the "hermes_state" logger name. logger = logging.getLogger("hermes_state") +# ASCII controls (keeping \t \n \r for the whitespace collapse), then zero-width, +# bidi override, object-replacement and interlinear-annotation code points. +_TITLE_CONTROL_RE = re.compile(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]') +_TITLE_INVISIBLE_RE = re.compile(r'[\u200b-\u200f\u2028-\u202e\u2060-\u2069\ufeff\ufffc\ufff9-\ufffb]') +_NUMBERED_TITLE_RE = re.compile(r'^(.*?) #(\d+)$') + class SessionTitlesMixin: """Sanitizing, ranking auto/user titles, lineage-aware lookups.""" @classmethod def _title_rank(cls, source: Optional[str]) -> int: - """Rank a stored title_source. - - NULL (pre-provenance rows) is indistinguishable from a manual ``/title`` - of that era, so it ranks as ``user``: auto-titling only ever fills - genuinely empty legacy titles. - """ + """Rank a stored title_source. NULL (pre-provenance rows) is indistinguishable + from a manual ``/title`` of that era, so it ranks as ``user``.""" if source is None: return cls._TITLE_SOURCE_RANK[cls.TITLE_SOURCE_USER] return cls._TITLE_SOURCE_RANK.get(str(source), 0) @staticmethod def sanitize_title(title: Optional[str]) -> Optional[str]: - """Strip control/zero-width/bidi chars, collapse whitespace, normalize - empty to None. Raises ValueError if longer than MAX_TITLE_LENGTH - after cleaning.""" + """Strip control/zero-width/bidi chars (and lone surrogates sqlite3 cannot + bind), collapse whitespace, normalize empty to None. Raises ValueError if + longer than MAX_TITLE_LENGTH after cleaning.""" from hermes_state import SessionDB if not title: return None - - # Lone surrogates cannot be bound by sqlite3 (UnicodeEncodeError). - title = _sanitize_surrogates(title) - - # ASCII controls, keeping \t \n \r so the whitespace collapse below - # turns them into spaces. - cleaned = re.sub(r'[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]', '', title) - - # Zero-width, bidi override, object-replacement, interlinear annotation. - cleaned = re.sub( - r'[\u200b-\u200f\u2028-\u202e\u2060-\u2069\ufeff\ufffc\ufff9-\ufffb]', - '', cleaned, - ) - + cleaned = _TITLE_INVISIBLE_RE.sub('', _TITLE_CONTROL_RE.sub('', _sanitize_surrogates(title))) cleaned = re.sub(r'\s+', ' ', cleaned).strip() - if not cleaned: return None - if len(cleaned) > SessionDB.MAX_TITLE_LENGTH: - raise ValueError( - f"Title too long ({len(cleaned)} chars, max {SessionDB.MAX_TITLE_LENGTH})" - ) - + raise ValueError(f"Title too long ({len(cleaned)} chars, max {SessionDB.MAX_TITLE_LENGTH})") return cleaned - def _is_compression_ancestor( - self, conn, *, ancestor_id: str, descendant_id: str - ) -> bool: - """True if *ancestor_id* is a compression predecessor of *descendant_id*. - - Uses the canonical continuation edge ``_COMPRESSION_CHILD_SQL`` (parent - ended with ``end_reason = 'compression'`` and child started at/after its - ``ended_at``), which excludes delegate/branch children that also carry - ``parent_session_id``. One recursive CTE so the edge is defined once. - """ + def _is_compression_ancestor(self, conn, *, ancestor_id: str, descendant_id: str) -> bool: + """True if *ancestor_id* is a compression predecessor of *descendant_id*, via the + canonical continuation edge ``_COMPRESSION_CHILD_SQL`` (excludes delegate/branch + children that also carry ``parent_session_id``).""" if not ancestor_id or not descendant_id or ancestor_id == descendant_id: return False edge = _COMPRESSION_CHILD_SQL.format(a="child") @@ -93,23 +71,15 @@ class SessionTitlesMixin: ).fetchone() return row is not None - def _set_session_title( - self, - session_id: str, - title: str, - *, - source: str, - ) -> bool: + def _set_session_title(self, session_id: str, title: str, *, source: str) -> bool: """Write a title, enforcing provenance precedence. - A ``user`` write always lands. ``derived``/``llm`` land only when the - row is untitled or holds strictly lower authority, so derived upgrades - to llm exactly once, nothing overwrites a user name, and re-running the - titler on an llm row is a no-op (stops sessions renaming themselves). - No writer may move a hidden canonical Bot Chat off its title. - - Read and write are one compare-and-swap in a single transaction, so a - manual ``/title`` racing an in-flight generation is not clobbered. + A ``user`` write always lands. ``derived``/``llm`` land only when the row is + untitled or holds strictly lower authority (derived upgrades to llm exactly once, + 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. """ title = self.sanitize_title(title) is_user = source == self.TITLE_SOURCE_USER @@ -117,18 +87,14 @@ class SessionTitlesMixin: def _do(conn): current = conn.execute( - "SELECT title, title_source, hidden FROM sessions WHERE id = ?", - (session_id,), + "SELECT title, title_source, hidden FROM sessions WHERE id = ?", (session_id,), ).fetchone() if current is None: return 0 - # The canonical Bot Chat's NAME is its identity: Bot Mode resolves it - # by exact-title lookup on every open, so a rename orphans the whole - # conversation (next open mints an empty replacement and UNIQUE(title) - # blocks renaming back). Refuse here, the single write path every - # surface funnels through. Hidden is the discriminator: canonical - # chats are born hidden; a visible session merely named "Bot Chat" - # stays renameable. Provenance-blind so the auto-titler no-ops too. + # The canonical Bot Chat's NAME is its identity (Bot Mode resolves it by + # exact-title lookup on every open), so a rename orphans the conversation. + # Hidden is the discriminator: canonical chats are born hidden; a visible + # session merely named "Bot Chat" stays renameable. Provenance-blind. if ( (current["title"] or "") == self.CANONICAL_BOT_CHAT_TITLE and bool(current["hidden"]) @@ -141,100 +107,64 @@ class SessionTitlesMixin: "To start fresh, create a new bot instead." ) return 0 - if not is_user and current["title"] is not None: - if self._title_rank(current["title_source"]) >= new_rank: - return 0 - + if not is_user and current["title"] is not None and self._title_rank(current["title_source"]) >= new_rank: + return 0 if title: - cursor = conn.execute( - "SELECT id FROM sessions WHERE title = ? AND id != ?", - (title, session_id), - ) - conflict = cursor.fetchone() + conflict = conn.execute( + "SELECT id FROM sessions WHERE title = ? AND id != ?", (title, session_id), + ).fetchone() if conflict: conflict_id = conflict["id"] - # If the conflicting holder is a hidden compressed ancestor - # of this continuation, the user cannot free the title, so - # transfer it onto the tip. Uniqueness and lineage are kept. - if self._is_compression_ancestor( - conn, ancestor_id=conflict_id, descendant_id=session_id - ): - conn.execute( - "UPDATE sessions SET title = NULL WHERE id = ?", - (conflict_id,), - ) + # A hidden compressed ancestor holding the title cannot be freed by + # the user, so transfer it onto the tip (uniqueness + lineage kept). + if self._is_compression_ancestor(conn, ancestor_id=conflict_id, descendant_id=session_id): + conn.execute("UPDATE sessions SET title = NULL WHERE id = ?", (conflict_id,)) else: - raise ValueError( - f"Title '{title}' is already in use by session {conflict_id}" - ) - # CAS on the values just read (``IS`` is NULL-safe): a concurrent - # write between the SELECT and here loses instead of being overwritten. + raise ValueError(f"Title '{title}' is already in use by session {conflict_id}") + # CAS on the values just read (``IS`` is NULL-safe): a concurrent write + # between the SELECT and here loses instead of being overwritten. cursor = conn.execute( "UPDATE sessions SET title = ?, title_source = ? " "WHERE id = ? AND title IS ? AND title_source IS ?", - ( - title, - source if title else None, - session_id, - current["title"], - current["title_source"], - ), + (title, source if title else None, session_id, current["title"], current["title_source"]), ) return cursor.rowcount - rowcount = self._execute_write(_do) - return rowcount > 0 + return self._execute_write(_do) > 0 def set_session_title(self, session_id: str, title: str) -> bool: - """Set a title on the user's behalf (``user`` provenance; auto-titling - never replaces it). Empty clears the title. Raises ValueError on a - title conflict or validation failure. Automatic callers must use - :meth:`set_auto_title`.""" - return self._set_session_title( - session_id, title, source=self.TITLE_SOURCE_USER - ) + """Set a title on the user's behalf (``user`` provenance). Empty clears it. + Raises ValueError on conflict or validation failure.""" + return self._set_session_title(session_id, title, source=self.TITLE_SOURCE_USER) def set_auto_title(self, session_id: str, title: str, *, source: str) -> bool: - """Set an automatic title; False (untouched) when a higher-authority - title already holds the row.""" + """Set an automatic title; False (untouched) when a higher-authority title + already holds the row.""" if source not in (self.TITLE_SOURCE_DERIVED, self.TITLE_SOURCE_LLM): raise ValueError(f"invalid automatic title source: {source!r}") return self._set_session_title(session_id, title, source=source) def set_auto_title_if_empty(self, session_id: str, title: str) -> bool: - """Back-compat shim (third-party plugins reference it by name); new - code calls :meth:`set_auto_title` with an explicit source.""" - return self.set_auto_title( - session_id, title, source=self.TITLE_SOURCE_LLM - ) + """Back-compat shim (third-party plugins reference it by name).""" + return self.set_auto_title(session_id, title, source=self.TITLE_SOURCE_LLM) def get_session_title(self, session_id: str) -> Optional[str]: """Get the title for a session, or None.""" - with self._read_ctx() as conn: - cursor = conn.execute( - "SELECT title FROM sessions WHERE id = ?", (session_id,) - ) - row = cursor.fetchone() + row = self._read_one("SELECT title FROM sessions WHERE id = ?", (session_id,)) return row["title"] if row else None def get_session_title_source(self, session_id: str) -> Optional[str]: """Get the provenance of a session's title, or None when untitled.""" - with self._read_ctx() as conn: - cursor = conn.execute( - "SELECT title, title_source FROM sessions WHERE id = ?", - (session_id,), - ) - row = cursor.fetchone() + row = self._read_one("SELECT title, title_source FROM sessions WHERE id = ?", (session_id,)) if not row or row["title"] is None: return None return row["title_source"] def set_session_title_source(self, session_id: str, source: str) -> bool: - """Overwrite a title's provenance without touching the text: a title - copied across a compression rotation keeps the original's authority.""" + """Overwrite a title's provenance without touching the text (a title copied + across a compression rotation keeps the original's authority).""" 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", @@ -243,63 +173,44 @@ class SessionTitlesMixin: def get_session_by_title(self, title: str) -> Optional[Dict[str, Any]]: """Look up a session by exact title. Returns session dict or None.""" - with self._read_ctx() as conn: - cursor = conn.execute( - "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,), - ) - row = cursor.fetchone() + 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,), + ) return self._session_row_dict(row) if row else None def resolve_session_by_title(self, title: str) -> Optional[str]: """Resolve a title to a session ID, preferring the latest "title #N" continuation over the exact match.""" exact = self.get_session_by_title(title) - # Escape LIKE wildcards so "%"/"_" in titles cannot false-match. - escaped = _escape_like(title) - with self._read_ctx() as conn: - cursor = conn.execute( - "SELECT id, title, started_at FROM sessions " - "WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC", - (f"{escaped} #%",), - ) - numbered = cursor.fetchall() - + numbered = self._read_all( + "SELECT id, title, started_at FROM sessions " + "WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC", + (f"{_escape_like(title)} #%",), + ) if numbered: return numbered[0]["id"] - elif exact: - return exact["id"] - return None + return exact["id"] if exact else None def get_next_title_in_lineage(self, base_title: str) -> str: - """Next title in a lineage ("my session" → "my session #2"): strip any - " #N" suffix, then increment the highest existing number.""" - match = re.match(r'^(.*?) #(\d+)$', base_title) - if match: - base = match.group(1) - else: - base = base_title - - escaped = _escape_like(base) - with self._read_ctx() as conn: - cursor = conn.execute( - "SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'", - (base, f"{escaped} #%"), - ) - existing = [row["title"] for row in cursor.fetchall()] - - if not existing: + """Next title in a lineage ("my session" -> "my session #2"): strip any " #N" + suffix, then increment the highest existing number.""" + match = _NUMBERED_TITLE_RE.match(base_title) + 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)} #%"), + ) + if not rows: return base - - max_num = 1 # The unnumbered original counts as #1 - for t in existing: - m = re.match(r'^.* #(\d+)$', t) + 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}" diff --git a/hermes_state_usage.py b/hermes_state_usage.py index 9abae52aa7..d889e6d106 100644 --- a/hermes_state_usage.py +++ b/hermes_state_usage.py @@ -14,24 +14,79 @@ from typing import Any, Dict, List, Optional, Tuple # caplog tests pin the "hermes_state" logger name. logger = logging.getLogger("hermes_state") +_TOKEN_UPDATE_ABSOLUTE_SQL = """UPDATE sessions SET + input_tokens = ?, + output_tokens = ?, + cache_read_tokens = ?, + cache_write_tokens = ?, + reasoning_tokens = ?, + estimated_cost_usd = COALESCE(?, 0), + actual_cost_usd = CASE + WHEN ? IS NULL THEN actual_cost_usd + ELSE ? + END, + cost_status = COALESCE(?, cost_status), + cost_source = COALESCE(?, cost_source), + pricing_version = COALESCE(?, pricing_version), + billing_provider = COALESCE(billing_provider, ?), + billing_base_url = COALESCE(billing_base_url, ?), + billing_mode = COALESCE(billing_mode, ?), + model = COALESCE(model, ?), + api_call_count = ? + WHERE id = ?""" + +_TOKEN_UPDATE_DELTA_SQL = """UPDATE sessions SET + 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 = COALESCE(estimated_cost_usd, 0) + COALESCE(?, 0), + actual_cost_usd = CASE + WHEN ? IS NULL THEN actual_cost_usd + ELSE COALESCE(actual_cost_usd, 0) + ? + END, + cost_status = COALESCE(?, cost_status), + cost_source = COALESCE(?, cost_source), + pricing_version = COALESCE(?, pricing_version), + billing_provider = COALESCE(billing_provider, ?), + billing_base_url = COALESCE(billing_base_url, ?), + billing_mode = COALESCE(billing_mode, ?), + model = COALESCE(model, ?), + api_call_count = COALESCE(api_call_count, 0) + ? + WHERE id = ?""" + +_MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage ( + session_id, model, billing_provider, billing_base_url, billing_mode, + task, api_call_count, input_tokens, output_tokens, + cache_read_tokens, cache_write_tokens, reasoning_tokens, + estimated_cost_usd, actual_cost_usd, cost_status, cost_source, + first_seen, last_seen + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(session_id, model, billing_provider, billing_base_url, billing_mode, task) + DO UPDATE SET + api_call_count = api_call_count + excluded.api_call_count, + input_tokens = input_tokens + excluded.input_tokens, + output_tokens = output_tokens + excluded.output_tokens, + cache_read_tokens = cache_read_tokens + excluded.cache_read_tokens, + cache_write_tokens = cache_write_tokens + excluded.cache_write_tokens, + reasoning_tokens = reasoning_tokens + excluded.reasoning_tokens, + estimated_cost_usd = estimated_cost_usd + excluded.estimated_cost_usd, + actual_cost_usd = actual_cost_usd + excluded.actual_cost_usd, + cost_status = COALESCE(excluded.cost_status, cost_status), + cost_source = COALESCE(excluded.cost_source, cost_source), + last_seen = excluded.last_seen""" + class SessionUsageMixin: """Coalesced token writer, per-model usage rows, billing route.""" def update_session_billing_route( - self, - session_id: str, - *, - provider: str, - base_url: str, - billing_mode: Optional[str] = None, + self, session_id: str, *, provider: str, base_url: str, billing_mode: Optional[str] = None, ) -> None: """Unconditionally set the billing route (``update_token_counts`` only - COALESCE-fills NULLs) so the dashboard reflects the latest /model switch. - - Also nulls ``system_prompt`` so the cached snapshot (stale ``Model:`` / - ``Provider:`` header) is rebuilt, like ``update_session_model``. - """ + COALESCE-fills NULLs) so the dashboard reflects the latest /model switch. Also + nulls ``system_prompt`` so the cached snapshot header is rebuilt.""" # Barrier against queued token deltas — see update_session_model. self.flush_token_counts() @@ -50,29 +105,20 @@ class SessionUsageMixin: self._execute_write(_do) def queue_token_counts(self, session_id: str, **kwargs) -> None: - """Enqueue a token/cost delta for the background writer. - - Same kwargs and semantics as :meth:`update_token_counts`, applied - asynchronously; cheap enough for the turn thread. After close() has - stopped the writer, falls back to the synchronous path and may raise. - """ + """Enqueue a token/cost delta for the background writer (same kwargs as + :meth:`update_token_counts`). After close() has stopped the writer, falls back + 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_stopped = self._token_writer_stop and (thread is None or not thread.is_alive()) if not writer_stopped: self._token_queue.append((session_id, kwargs)) if thread is None or not thread.is_alive(): - # Daemon so exit never hangs on accounting; the atexit hook - # (registered once per instance) drains leftovers. Checking - # ``not is_alive()`` rather than ``is None`` respawns a writer - # that died from an unexpected escape, otherwise deltas - # would pile up until a reader's flush drained them. + # 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() @@ -88,18 +134,21 @@ class SessionUsageMixin: atexit.register(_drain_at_exit) self._token_queue_cond.notify_all() if writer_stopped: - # close() ran (a stop-flagged but live writer still accepts; its - # loop drains before exiting). Enqueueing now would drop the delta - # silently — no writer, atexit hook gone — so apply inline and let a - # closed-connection failure raise at the call site. + # close() ran: enqueueing would drop the delta silently, so apply inline. self.update_token_counts(session_id, **kwargs) - def flush_token_counts(self, timeout: float = 5.0) -> bool: - """Block until every queued token delta has been applied. + def _apply_claimed_batch(self, batch) -> None: + """Apply a batch whose ``busy`` flag the caller already claimed, then release.""" + try: + self._apply_token_batch(batch) + finally: + with self._token_queue_cond: + self._token_writer_busy = False + self._token_queue_cond.notify_all() - False on timeout (callers then read totals stale by the queued deltas). - Never raises: apply failures are logged by the writer. - """ + def flush_token_counts(self, timeout: float = 5.0) -> bool: + """Block until every queued token delta has been applied. False on timeout + (callers then read totals stale by the queued deltas). Never raises.""" # Lock-free fast path: reads queue-then-busy (see ordering notes below). if not self._token_queue and not self._token_writer_busy: return True @@ -107,20 +156,13 @@ class SessionUsageMixin: with self._token_queue_cond: deadline = time.monotonic() + timeout while self._token_queue or self._token_writer_busy: - # A live writer is authoritative even when stop-flagged: draining - # here would race its in-flight batch, and newer deltas committing - # before older ones breaks last-non-None-wins / first-accounted- - # route / COALESCE-backfill fields. Only a dead writer lets the - # caller take leftovers; re-checked each wakeup because the writer - # can exit mid-wait with deltas enqueued after its final check. - # busy is claimed while draining so a concurrent flush cannot - # report drained or pop a newer delta while this batch is - # unapplied: a claimed busy means "wait", never "drain alongside". + # A live writer is authoritative even when stop-flagged: draining here + # would race its in-flight batch and reorder deltas (breaking last-non- + # None-wins / first-accounted-route / COALESCE-backfill fields). Only a + # dead writer lets the caller take leftovers; a claimed busy means + # "wait", never "drain alongside". thread = self._token_writer_thread - if ( - (thread is None or not thread.is_alive()) - and not self._token_writer_busy - ): + if (thread is None or not thread.is_alive()) and not self._token_writer_busy: self._token_writer_busy = True batch = list(self._token_queue) self._token_queue.clear() @@ -130,12 +172,7 @@ class SessionUsageMixin: return False self._token_queue_cond.wait(remaining) if batch: - try: - self._apply_token_batch(batch) - finally: - with self._token_queue_cond: - self._token_writer_busy = False - self._token_queue_cond.notify_all() + self._apply_claimed_batch(batch) return True def _token_writer_loop(self) -> None: @@ -145,62 +182,44 @@ class SessionUsageMixin: while not self._token_queue and not self._token_writer_stop: remaining = idle_deadline - time.monotonic() if remaining <= 0: - # Retire under the same lock queue_token_counts() uses to - # decide to spawn, so no delta strands behind an exiting worker. + # Retire under the same lock queue_token_counts() uses to decide + # to spawn, so no delta strands behind an exiting worker. self._token_writer_thread = None return self._token_queue_cond.wait(remaining) if not self._token_queue: self._token_writer_thread = None return # stop requested and fully drained - # busy BEFORE clearing the queue: flush's lock-free fast path - # reads queue-then-busy and must never see "empty and idle" - # while a popped batch is unapplied. + # busy BEFORE clearing the queue: flush's lock-free fast path must never + # see "empty and idle" while a popped batch is unapplied. self._token_writer_busy = True batch = list(self._token_queue) self._token_queue.clear() - try: - self._apply_token_batch(batch) - finally: - with self._token_queue_cond: - self._token_writer_busy = False - self._token_queue_cond.notify_all() + self._apply_claimed_batch(batch) def _apply_token_batch(self, batch: List[Tuple[str, Dict[str, Any]]]) -> None: """Apply queued deltas in order, coalescing where safe. Never raises.""" try: coalesced = self._coalesce_token_deltas(batch) except Exception as exc: - # Coalescing must never kill the writer (callers cannot observe a - # dead one); the merge is only an optimization. - logger.warning( - "async token accounting: coalesce failed, applying raw " - "batch: %s", exc, - ) + # Coalescing must never kill the writer; the merge is only an optimization. + logger.warning("async token accounting: coalesce failed, applying raw batch: %s", exc) coalesced = batch for session_id, kwargs in coalesced: try: self.update_token_counts(session_id, **kwargs) except Exception as exc: # Accounting loss is logged, never raised into a turn. - logger.warning( - "async token accounting: apply failed (session=%s): %s", - session_id, exc, - ) + logger.warning("async token accounting: apply failed (session=%s): %s", session_id, exc) - def _coalesce_token_deltas( - self, batch: List[Tuple[str, Dict[str, Any]]] - ) -> List[Tuple[str, Dict[str, Any]]]: - """Merge adjacent incremental deltas with an identical route, so - ordering across sessions and /model switches is preserved exactly. - absolute=True deltas never merge.""" + def _coalesce_token_deltas(self, batch: List[Tuple[str, Dict[str, Any]]]) -> List[Tuple[str, Dict[str, Any]]]: + """Merge adjacent incremental deltas with an identical route, so ordering across + sessions and /model switches is preserved exactly. absolute=True never merges.""" groups: List[Tuple[Optional[tuple], str, Dict[str, Any]]] = [] for session_id, kwargs in batch: key = None if not kwargs.get("absolute"): - key = (session_id,) + tuple( - kwargs.get(f) for f in self._TOKEN_DELTA_ROUTE_FIELDS - ) + key = (session_id,) + tuple(kwargs.get(f) for f in self._TOKEN_DELTA_ROUTE_FIELDS) if groups and key is not None and groups[-1][0] == key: merged = groups[-1][2] for f in self._TOKEN_DELTA_SUM_FIELDS: @@ -223,18 +242,16 @@ class SessionUsageMixin: if thread is not None and thread.is_alive(): thread.join(timeout=join_timeout) if thread.is_alive(): - # Writer stuck mid-apply: leave deltas unapplied rather than - # race it and misorder/double-count. + # 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), ) return - # Writer gone: apply leftovers synchronously under the same busy - # protocol. Wait out a flush caller-drain that already claimed busy — - # close() nulls the connection right after this returns and must not - # yank it mid-batch. + # Writer gone: apply leftovers synchronously under the same busy protocol. Wait + # out a flush caller-drain that already claimed busy — close() nulls the + # connection right after this returns and must not yank it mid-batch. with self._token_queue_cond: deadline = time.monotonic() + join_timeout while self._token_writer_busy: @@ -247,19 +264,13 @@ class SessionUsageMixin: ) return self._token_queue_cond.wait(remaining) - # busy BEFORE clearing the queue (same ordering as the writer loop), - # or flush's lock-free fast path could see "empty and idle". + # busy BEFORE clearing the queue (same ordering as the writer loop). batch = list(self._token_queue) if batch: self._token_writer_busy = True self._token_queue.clear() if batch: - try: - self._apply_token_batch(batch) - finally: - with self._token_queue_cond: - self._token_writer_busy = False - self._token_queue_cond.notify_all() + self._apply_claimed_batch(batch) def _drain_token_queue_at_exit(self) -> None: try: @@ -287,75 +298,21 @@ class SessionUsageMixin: api_call_count: int = 0, absolute: bool = False, ) -> None: - """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). - """ - # Ensure the row exists: under concurrent load the initial - # create_session() may have failed on SQLite locking, and the UPDATE - # would silently affect 0 rows. + """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).""" + # 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) - if absolute: - sql = """UPDATE sessions SET - input_tokens = ?, - output_tokens = ?, - cache_read_tokens = ?, - cache_write_tokens = ?, - reasoning_tokens = ?, - estimated_cost_usd = COALESCE(?, 0), - actual_cost_usd = CASE - WHEN ? IS NULL THEN actual_cost_usd - ELSE ? - END, - cost_status = COALESCE(?, cost_status), - cost_source = COALESCE(?, cost_source), - pricing_version = COALESCE(?, pricing_version), - billing_provider = COALESCE(billing_provider, ?), - billing_base_url = COALESCE(billing_base_url, ?), - billing_mode = COALESCE(billing_mode, ?), - model = COALESCE(model, ?), - api_call_count = ? - WHERE id = ?""" - else: - sql = """UPDATE sessions SET - 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 = COALESCE(estimated_cost_usd, 0) + COALESCE(?, 0), - actual_cost_usd = CASE - WHEN ? IS NULL THEN actual_cost_usd - ELSE COALESCE(actual_cost_usd, 0) + ? - END, - cost_status = COALESCE(?, cost_status), - cost_source = COALESCE(?, cost_source), - pricing_version = COALESCE(?, pricing_version), - billing_provider = COALESCE(billing_provider, ?), - billing_base_url = COALESCE(billing_base_url, ?), - billing_mode = COALESCE(billing_mode, ?), - model = COALESCE(model, ?), - api_call_count = COALESCE(api_call_count, 0) + ? - WHERE id = ?""" - has_accounted_usage = bool( + sql = _TOKEN_UPDATE_ABSOLUTE_SQL if absolute else _TOKEN_UPDATE_DELTA_SQL + has_usage = bool( input_tokens or output_tokens or cache_read_tokens - or cache_write_tokens or reasoning_tokens or api_call_count - or estimated_cost_usd or actual_cost_usd + or cache_write_tokens or reasoning_tokens or api_call_count or estimated_cost_usd ) + has_accounted_usage = bool(has_usage or actual_cost_usd) params = ( - input_tokens, - output_tokens, - cache_read_tokens, - cache_write_tokens, - reasoning_tokens, - estimated_cost_usd, - actual_cost_usd, - actual_cost_usd, - cost_status, - cost_source, - pricing_version, + input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens, + estimated_cost_usd, actual_cost_usd, actual_cost_usd, cost_status, cost_source, pricing_version, 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, @@ -363,31 +320,22 @@ class SessionUsageMixin: 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. Each delta carries the route active at call time - # and is recorded into session_model_usage keyed by it. Only the - # incremental path records here: absolute cumulative updates cannot be + # 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 # split back into routes; Insights reconciles the residual instead. - record_model_usage = (not absolute) and ( - input_tokens or output_tokens or cache_read_tokens - or cache_write_tokens or reasoning_tokens or api_call_count - or estimated_cost_usd - ) + record_model_usage = (not absolute) and has_usage def _do(conn): row = conn.execute( - "SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", - (session_id,), + "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) - - # 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-provider usage). + # 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 @@ -406,21 +354,12 @@ 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, + 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._execute_write(_do) @@ -446,77 +385,31 @@ class SessionUsageMixin: api_call_count: int, task: str = "", ) -> None: - """Accumulate a per-API-call usage delta into session_model_usage. - - Runs inside the caller's write transaction, after the ``sessions`` - UPDATE, so per-model rows stay consistent with the summary row. A - missing model/provider falls back to the session row (same COALESCE - behaviour as the summary update). ``task`` is ``''`` for the main loop; - auxiliary calls record their task name via :meth:`record_auxiliary_usage`. + """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. """ row = conn.execute( "SELECT model, billing_provider, billing_base_url, billing_mode " "FROM sessions WHERE id = ?", (session_id,), ).fetchone() - sess_model = row["model"] if row is not None else None - sess_provider = row["billing_provider"] if row is not None else None - sess_base_url = row["billing_base_url"] if row is not None else None - sess_billing_mode = row["billing_mode"] if row is not None else None - - # Aux rows must NOT inherit the main-loop route (vision on gemini while - # the main loop runs anthropic); missing info stays 'unknown'/empty. - if task: - eff_model = model or "unknown" - eff_provider = billing_provider or "" - eff_base_url = billing_base_url or "" - eff_billing_mode = billing_mode or "" - else: - eff_model = model or sess_model or "unknown" - eff_provider = billing_provider or sess_provider or "" - eff_base_url = billing_base_url or sess_base_url or "" - eff_billing_mode = billing_mode or sess_billing_mode or "" + 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( - """INSERT INTO session_model_usage ( - session_id, model, billing_provider, billing_base_url, billing_mode, - task, api_call_count, input_tokens, output_tokens, - cache_read_tokens, cache_write_tokens, reasoning_tokens, - estimated_cost_usd, actual_cost_usd, cost_status, cost_source, - first_seen, last_seen - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) - ON CONFLICT(session_id, model, billing_provider, billing_base_url, billing_mode, task) - DO UPDATE SET - api_call_count = api_call_count + excluded.api_call_count, - input_tokens = input_tokens + excluded.input_tokens, - output_tokens = output_tokens + excluded.output_tokens, - cache_read_tokens = cache_read_tokens + excluded.cache_read_tokens, - cache_write_tokens = cache_write_tokens + excluded.cache_write_tokens, - reasoning_tokens = reasoning_tokens + excluded.reasoning_tokens, - estimated_cost_usd = estimated_cost_usd + excluded.estimated_cost_usd, - actual_cost_usd = actual_cost_usd + excluded.actual_cost_usd, - cost_status = COALESCE(excluded.cost_status, cost_status), - cost_source = COALESCE(excluded.cost_source, cost_source), - last_seen = excluded.last_seen""", + _MODEL_USAGE_UPSERT_SQL, ( - session_id, - eff_model, - eff_provider, - eff_base_url, - eff_billing_mode, - task or "", - api_call_count or 0, - input_tokens or 0, - output_tokens or 0, - cache_read_tokens or 0, - cache_write_tokens or 0, - reasoning_tokens or 0, - float(estimated_cost_usd or 0.0), - float(actual_cost_usd or 0.0), - cost_status, - cost_source, - now, - now, + session_id, eff_model, eff_provider, eff_base_url, eff_billing_mode, 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, ), ) @@ -536,16 +429,11 @@ class SessionUsageMixin: estimated_cost_usd: Optional[float] = None, api_call_count: int = 1, ) -> None: - """Record an auxiliary LLM call's usage (vision, compression, title - generation, ...) against *session_id*. - - Writes a per-(model, provider, task) delta into ``session_model_usage`` - WITHOUT touching the ``sessions`` summary row: the gateway overwrites - session counters with absolute main-loop totals, so aux tokens there - would be clobbered or double-counted. Insights read the union. - ``api_call_count`` may aggregate N calls (background-review forks). - Best-effort: callers must never fail an aux call over accounting. - """ + """Record an auxiliary LLM call's usage (vision, compression, title generation, + ...) as a per-(model, provider, task) delta in ``session_model_usage`` WITHOUT + 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.""" if not session_id or not task: return # FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts. @@ -553,37 +441,24 @@ class SessionUsageMixin: 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) - ), + 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) 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 not shrink with paging. Spend prefers the billed figure over - the estimate, the same precedence a single row renders.""" + """Tokens and spend across the whole store (one scan), so the sidebar total does + not shrink with paging. Spend prefers the billed figure over the estimate.""" where = ["parent_session_id IS NULL", "message_count >= ?"] params: List[Any] = [min_message_count] if not include_archived: where.append("COALESCE(archived, 0) = 0") - row = self._read_one( f""" SELECT COALESCE(SUM(COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0)), 0), @@ -593,5 +468,4 @@ class SessionUsageMixin: """, params, ) - return {"tokens": int(row[0] or 0), "cost_usd": float(row[1] or 0.0)}