fix(state): isolate background reads from writer connection

This commit is contained in:
milnerrad
2026-08-09 18:12:35 +08:00
committed by Teknium
parent d16622cd14
commit 94869e5a0d
2 changed files with 59 additions and 28 deletions

View File

@@ -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 []

View File

@@ -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")