diff --git a/hermes_state.py b/hermes_state.py index 973e67d1a3..c5f5aad888 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -55,7 +55,7 @@ from hermes_constants import get_hermes_home from hermes_cli.sqlite_runtime import ( is_sqlite_wal_reset_vulnerable as _is_sqlite_wal_reset_vulnerable, ) -from typing import Any, Callable, Dict, List, Optional, Set, Tuple, TypeVar +from typing import Any, Callable, Dict, Iterator, List, Optional, Set, Tuple, TypeVar, cast from hermes_state_common import ( # noqa: F401 (re-exported for back-compat) _BRANCH_CHILD_SQL, @@ -4969,7 +4969,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) return self._get_read_conn() @contextmanager - def _read_ctx(self): + def _read_ctx(self) -> Iterator[sqlite3.Connection]: """Yield a connection for read-only statements. WAL: a read-only connection borrowed from a bounded pool with NO @@ -5012,7 +5012,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) self._close_read_conn(conn) return with self._lock: - yield self._conn + yield cast(sqlite3.Connection, self._conn) # ── Core write helper ── @@ -8201,11 +8201,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) if not session_id: return None now = time.time() - row = self._conn.execute( - "SELECT holder FROM compression_locks " - "WHERE session_id = ? AND expires_at >= ?", - (session_id, now), - ).fetchone() + with self._read_ctx() as conn: + row = conn.execute( + "SELECT holder FROM compression_locks " + "WHERE session_id = ? AND expires_at >= ?", + (session_id, now), + ).fetchone() if row is None: return None return row["holder"] if isinstance(row, sqlite3.Row) else row[0] @@ -8279,11 +8280,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # No-op fast path: skip the transaction when there is nothing to # clear. Read-only, no write lock. try: - row = self._conn.execute( - "SELECT last_activity_description, last_activity_provenance " - "FROM sessions WHERE id = ?", - (session_id,), - ).fetchone() + with self._read_ctx() as conn: + row = conn.execute( + "SELECT last_activity_description, last_activity_provenance " + "FROM sessions WHERE id = ?", + (session_id,), + ).fetchone() except sqlite3.Error: row = None if row is not None: @@ -15244,12 +15246,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) no handoff record. """ try: - cur = self._conn.execute( - "SELECT handoff_state, handoff_platform, handoff_error " - "FROM sessions WHERE id = ?", - (session_id,), - ) - row = cur.fetchone() + with self._read_ctx() as conn: + row = conn.execute( + "SELECT handoff_state, handoff_platform, handoff_error " + "FROM sessions WHERE id = ?", + (session_id,), + ).fetchone() if not row: return None return { @@ -15266,15 +15268,16 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) Used by the gateway's handoff watcher. """ try: - cur = self._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.handoff_state = 'pending' " - "ORDER BY s.started_at ASC" - ) - return [self._session_row_dict(r) for r in cur.fetchall()] + with self._read_ctx() as conn: + rows = 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.handoff_state = 'pending' " + "ORDER BY s.started_at ASC" + ).fetchall() + return [self._session_row_dict(r) for r in rows] except Exception: return [] diff --git a/tests/test_session_db_read_path_split.py b/tests/test_session_db_read_path_split.py index 1f3095ca87..67e57abb67 100644 --- a/tests/test_session_db_read_path_split.py +++ b/tests/test_session_db_read_path_split.py @@ -81,6 +81,34 @@ def test_reads_do_not_take_writer_lock(db): +@pytest.mark.requires_wal +def test_background_state_reads_never_touch_shared_writer(db): + """Background pollers must not race transcript writes on ``_conn``.""" + db.try_acquire_compression_lock("s1", "holder", ttl_seconds=60) + assert db.request_handoff("s1", "telegram") is True + # Open this thread's read connection before poisoning the shared writer. + assert db._get_read_conn() is not None + + writer = db._conn + + class PoisonWriter: + def execute(self, *_args, **_kwargs): + raise AssertionError("read touched shared writer connection") + + db._conn = PoisonWriter() + try: + assert db.get_compression_lock_holder("s1") == "holder" + db.clear_session_activity_labels("s1") + assert db.get_handoff_state("s1") == { + "state": "pending", + "platform": "telegram", + "error": None, + } + assert [row["id"] for row in db.list_pending_handoffs()] == ["s1"] + finally: + db._conn = writer + + def test_read_your_writes(db): """A fresh committed write must be visible to the read connection.""" db.append_message("s1", role="user", content="zanzibar checkpoint")