From 4de8710b745c8b1eed850474394888929ace11ea Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 19:43:49 -0700 Subject: [PATCH] refactor(state): unify _placeholders/_ended_by_compression/row-probe SQL into hermes_state_common --- hermes_state_common.py | 14 ++++++++++++++ hermes_state_compression.py | 11 ++++------- hermes_state_maintenance.py | 14 +++++--------- hermes_state_messages.py | 19 ++++++------------- 4 files changed, 29 insertions(+), 29 deletions(-) diff --git a/hermes_state_common.py b/hermes_state_common.py index 0f31877d8f..df39545504 100644 --- a/hermes_state_common.py +++ b/hermes_state_common.py @@ -292,6 +292,20 @@ def stat_db_file_identity(path) -> "tuple[int, int] | None": return (st.st_dev, st.st_ino) if st.st_dev and st.st_ino else None +# Row probes shared by the messages / compression mixins. +_ENDED_ROW_SQL = "SELECT ended_at, end_reason FROM sessions WHERE id = ?" +_COMPRESSION_LOCK_ROW_SQL = "SELECT holder, expires_at FROM compression_locks WHERE session_id = ?" + + +def _ended_by_compression(row) -> bool: + return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" + + +def _placeholders(items) -> str: + """``?,?,?`` for one bound parameter per element of *items* (a sequence or an int count).""" + return ",".join("?" for _ in range(items if isinstance(items, int) else len(items))) + + _FTS_TRIGGERS = ( "messages_fts_insert", "messages_fts_delete", "messages_fts_update", "messages_fts_trigram_insert", "messages_fts_trigram_delete", "messages_fts_trigram_update", diff --git a/hermes_state_compression.py b/hermes_state_compression.py index 7578fc73d8..9b57fd532e 100644 --- a/hermes_state_compression.py +++ b/hermes_state_compression.py @@ -11,22 +11,19 @@ import sqlite3 import time from typing import Any, Dict, List, Optional, Tuple -from hermes_state_common import _sql_session_last_active, is_automatic_end_reason +from hermes_state_common import ( + _COMPRESSION_LOCK_ROW_SQL as _LOCK_ROW_SQL, _ENDED_ROW_SQL, _ended_by_compression, _sql_session_last_active, + is_automatic_end_reason, +) # 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" - - def _cooldown_row(exists: bool, cooldown_until, error) -> Dict[str, Any]: return {"session_exists": exists, "cooldown_until": float(cooldown_until) if cooldown_until is not None else None, "error": error} diff --git a/hermes_state_maintenance.py b/hermes_state_maintenance.py index 22499b2c5c..69bc9e813e 100644 --- a/hermes_state_maintenance.py +++ b/hermes_state_maintenance.py @@ -8,7 +8,7 @@ from pathlib import Path from typing import Any, Dict, List, Optional, Tuple from hermes_state_common import ( - AUTO_VACUUM_MIN_FREELIST_RATIO, _sql_session_last_active, escape_like as _escape_like + AUTO_VACUUM_MIN_FREELIST_RATIO, _placeholders, _sql_session_last_active, escape_like as _escape_like ) # caplog tests pin the "hermes_state" logger name. @@ -37,10 +37,6 @@ def _one(clause: str, conv=None): return lambda v: ([clause], [conv(v) if conv else v]) -def _placeholders(n: int) -> str: - return ",".join("?" * n) - - def _seconds_since(now: float, raw) -> Optional[float]: """Age of a state_meta timestamp; None when unset or corrupt (= no prior run).""" try: @@ -100,7 +96,7 @@ class SessionMaintenanceMixin: ) """, (cutoff,)).fetchall()] if ids: - conn.execute(f"DELETE FROM sessions WHERE id IN ({_placeholders(len(ids))})", ids) + conn.execute(f"DELETE FROM sessions WHERE id IN ({_placeholders(ids)})", ids) self._delete_unreferenced_system_prompts(conn) return ids removed_ids = self._execute_write(_do) or [] @@ -159,7 +155,7 @@ class SessionMaintenanceMixin: orphan_predicate += (" AND NOT EXISTS (SELECT 1 FROM gateway_heartbeats h WHERE" " h.last_heartbeat >= ? AND h.started_at <= sessions.started_at + ?)") heartbeat_params = (now - hb_staleness, hb_grace) - scope_sql = f" AND source IN ({_placeholders(len(srcs))}){pin_scope} AND {orphan_predicate}" + scope_sql = f" AND source IN ({_placeholders(srcs)}){pin_scope} AND {orphan_predicate}" scope_params = (*srcs, cutoff, cutoff, *heartbeat_params) def _do(conn): rows = conn.execute(f"SELECT id FROM sessions WHERE ended_at IS NULL{scope_sql}", @@ -172,7 +168,7 @@ class SessionMaintenanceMixin: # Re-apply every predicate under the write lock. conn.execute( f"UPDATE sessions SET ended_at = ?, end_reason = 'startup_orphan_reap'" - f" WHERE id IN ({_placeholders(len(victims))}) AND ended_at IS NULL{scope_sql}", + f" WHERE id IN ({_placeholders(victims)}) AND ended_at IS NULL{scope_sql}", (time.time(), *victims, *scope_params)) return victims return self._execute_write(_do) or [] @@ -284,7 +280,7 @@ class SessionMaintenanceMixin: if not session_ids: return 0 conn.execute(f"UPDATE sessions SET parent_session_id = NULL " - f"WHERE parent_session_id IN ({_placeholders(len(session_ids))})", list(session_ids)) + f"WHERE parent_session_id IN ({_placeholders(session_ids)})", list(session_ids)) for sid in session_ids: conn.execute("DELETE FROM messages WHERE session_id = ?", (sid,)) conn.execute("DELETE FROM sessions WHERE id = ?", (sid,)) diff --git a/hermes_state_messages.py b/hermes_state_messages.py index 65a640d683..38a299716e 100644 --- a/hermes_state_messages.py +++ b/hermes_state_messages.py @@ -15,7 +15,10 @@ 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 ( + _COMPRESSION_LOCK_ROW_SQL, _ENDED_ROW_SQL, _RESET_END_REASONS, _RESET_END_REASONS_SQL, _ended_by_compression, + _legacy_reset_child_sql, _placeholders, +) # Log-record parity with the origin module (caplog tests pin "hermes_state"). logger = logging.getLogger("hermes_state") @@ -33,8 +36,6 @@ _BUMP_GENERATION_SQL = """ SET generation = conversation_generations.generation + 1 """ -_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 = ?" _DELETE_COMPRESSION_LOCK_SQL = "DELETE FROM compression_locks WHERE session_id = ? AND holder = ?" _DISPLAY_ACTIVE_CLAUSE = " AND (active = 1 OR compacted = 1)" @@ -46,10 +47,6 @@ _SET_DISPLAY_META_SQL = "UPDATE messages SET display_metadata = ? WHERE id = ?" _ARCHIVE_ACTIVE_SQL = "UPDATE messages SET active = 0, compacted = 1 WHERE session_id = ? AND active = 1" -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: @@ -97,10 +94,6 @@ def _tool_calls_count(tool_calls: Any) -> int: return 0 if tool_calls is None else (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" - - def _stale_holder(row, now: float) -> bool: """A lock/lease row whose holder is expired or a provably dead local process.""" from hermes_state import _compression_lock_holder_process_is_dead @@ -240,7 +233,7 @@ class SessionMessagesMixin: conn.execute( "DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?", (conversation_id, lease["holder"])) - session = conn.execute(_ENDED_BY_COMPRESSION_SQL, (session_id,)).fetchone() + session = conn.execute(_ENDED_ROW_SQL, (session_id,)).fetchone() if _ended_by_compression(session) and not allow_closed_compression_parent: raise CompressionSessionClosedError(session_id) @@ -507,7 +500,7 @@ class SessionMessagesMixin: if reject_active_turn_lease: self._check_transcript_write_guards( conn, session_id, None, reject_active_turn_lease=True, reject_active_compression_lock=True) - elif _ended_by_compression(conn.execute(_ENDED_BY_COMPRESSION_SQL, (session_id,)).fetchone()): + elif _ended_by_compression(conn.execute(_ENDED_ROW_SQL, (session_id,)).fetchone()): raise CompressionSessionClosedError(session_id) if archive_dropped: # Content-preserving UPDATE: FTS triggers don't fire on `active`, so the