refactor(state): drop blank separators around nested _do txn closures
This commit is contained in:
@@ -1012,7 +1012,6 @@ def fts_rebuild_admission(db_path, *, timeout_seconds=None):
|
||||
lock_path, exc)
|
||||
yield False
|
||||
return
|
||||
|
||||
acquired = False
|
||||
try:
|
||||
if _IS_WINDOWS:
|
||||
|
||||
@@ -82,7 +82,6 @@ class SessionCompressionMixin:
|
||||
an active lease or any canonical child means another path owns the lineage."""
|
||||
if not session_id:
|
||||
return False
|
||||
|
||||
def _do(conn):
|
||||
if not _ended_by_compression(conn.execute(_ENDED_ROW_SQL, (session_id,)).fetchone()):
|
||||
return False
|
||||
@@ -123,7 +122,6 @@ class SessionCompressionMixin:
|
||||
# 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,
|
||||
@@ -169,7 +167,6 @@ class SessionCompressionMixin:
|
||||
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(
|
||||
@@ -230,7 +227,6 @@ class SessionCompressionMixin:
|
||||
"WHERE id = ? AND ended_at IS NULL", (time.time(), parent_session_id))
|
||||
if updated.rowcount != 1:
|
||||
raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}")
|
||||
|
||||
self._execute_write(_do)
|
||||
|
||||
def _write_sql_logged(self, op: str, session_id: str, sql: str, params) -> None:
|
||||
@@ -279,7 +275,6 @@ class SessionCompressionMixin:
|
||||
return
|
||||
deadline = snapshot.get("cooldown_until")
|
||||
error = snapshot.get("error")
|
||||
|
||||
def _do(conn):
|
||||
cursor = conn.execute(
|
||||
"UPDATE sessions SET compression_failure_cooldown_until = ?, "
|
||||
@@ -385,7 +380,6 @@ class SessionCompressionMixin:
|
||||
return False
|
||||
now = time.time()
|
||||
expires_at = now + ttl_seconds
|
||||
|
||||
def _do(conn):
|
||||
return _claim_lease_row(
|
||||
conn, "compression_locks", "session_id", session_id, holder, now, expires_at,
|
||||
@@ -418,7 +412,6 @@ class SessionCompressionMixin:
|
||||
errors propagate so ``_execute_write`` can retry."""
|
||||
if not session_id:
|
||||
return session_id
|
||||
|
||||
def _row(sid: str):
|
||||
row = conn.execute(
|
||||
"SELECT id, parent_session_id, source, model_config, end_reason FROM sessions WHERE id = ?",
|
||||
@@ -457,14 +450,12 @@ class SessionCompressionMixin:
|
||||
return False
|
||||
now = time.time()
|
||||
expires_at = now + max(0.1, float(ttl_seconds))
|
||||
|
||||
def _do(conn):
|
||||
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
||||
return _claim_lease_row(
|
||||
conn, "session_turn_leases", "conversation_id", conversation_id, holder, now, expires_at,
|
||||
lambda h, e: float(e) <= now or _compression_lock_holder_process_is_dead(h),
|
||||
)[0]
|
||||
|
||||
return bool(self._execute_write(_do, patience_s=patience_s))
|
||||
|
||||
def acquire_session_turn_lease(
|
||||
@@ -517,27 +508,23 @@ class SessionCompressionMixin:
|
||||
if not session_id or not holder:
|
||||
return False
|
||||
expires_at = time.time() + max(0.1, float(ttl_seconds))
|
||||
|
||||
def _do(conn):
|
||||
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
||||
return conn.execute(
|
||||
"UPDATE session_turn_leases SET expires_at = ? "
|
||||
"WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder),
|
||||
).rowcount > 0
|
||||
|
||||
return bool(self._execute_write(_do))
|
||||
|
||||
def release_session_turn_lease(self, session_id: str, holder: str) -> None:
|
||||
"""Release a turn lease iff ``holder`` still owns it; idempotent."""
|
||||
if not session_id or not holder:
|
||||
return
|
||||
|
||||
def _do(conn):
|
||||
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
||||
conn.execute(
|
||||
"DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?",
|
||||
(conversation_id, holder))
|
||||
|
||||
self._execute_write(_do)
|
||||
|
||||
def get_compression_lock_holder(self, session_id: str) -> Optional[str]:
|
||||
|
||||
@@ -202,7 +202,6 @@ def quarantine_cross_process_lock(path: Path, timeout: float = 5.0):
|
||||
try:
|
||||
if platform.system() == "Windows":
|
||||
import msvcrt
|
||||
|
||||
def _lock(mode): # msvcrt locks a byte range from the current position
|
||||
handle.seek(0)
|
||||
msvcrt.locking(handle.fileno(), mode, 1)
|
||||
@@ -304,18 +303,15 @@ def collect_state_db_stats(db_path: Path) -> Dict[str, Any]:
|
||||
except Exception as exc:
|
||||
logger.debug("collect_state_db_stats: cannot open %s read-only: %s", db_path, exc)
|
||||
return stats
|
||||
|
||||
def _scalar(sql: str, params=()) -> Any:
|
||||
try:
|
||||
row = conn.execute(sql, params).fetchone()
|
||||
return row[0] if row else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _int(sql: str, params=()) -> Optional[int]:
|
||||
value = _scalar(sql, params)
|
||||
return int(value) if value is not None else None
|
||||
|
||||
def _meta_int(key: str) -> Optional[int]:
|
||||
try: # a non-numeric meta value must yield None, not fail the snapshot
|
||||
return _int("SELECT value FROM state_meta WHERE key = ?", (key,))
|
||||
|
||||
@@ -213,7 +213,6 @@ class SessionGatewayMixin:
|
||||
identity = (session_key, source, user_id, chat_id, chat_type, thread_id, display_name, origin_json)
|
||||
ancestors = include_compression_ancestors
|
||||
query_params = [session_id, *identity] if ancestors else [*identity, session_id]
|
||||
|
||||
def _do(conn):
|
||||
conn.execute(
|
||||
f"""{_COMPRESSION_LINEAGE_CTE if ancestors else ""}
|
||||
@@ -248,7 +247,6 @@ class SessionGatewayMixin:
|
||||
(session_id, source, user_id, session_key, chat_id, chat_type, thread_id, display_name,
|
||||
origin_json, self._own_profile_name(), time.time()),
|
||||
)
|
||||
|
||||
self._execute_write(_do)
|
||||
|
||||
def save_gateway_routing_entry(self, session_key: str, entry_json: str, *, scope: str = "") -> None:
|
||||
@@ -269,7 +267,6 @@ class SessionGatewayMixin:
|
||||
"""Atomically replace the routing index for *scope* (keys absent from *entries*
|
||||
are removed); other scopes untouched."""
|
||||
now = time.time()
|
||||
|
||||
def _do(conn):
|
||||
conn.execute("DELETE FROM gateway_routing WHERE scope = ?", (scope,))
|
||||
if entries:
|
||||
@@ -277,7 +274,6 @@ class SessionGatewayMixin:
|
||||
"INSERT INTO gateway_routing (scope, session_key, entry_json, updated_at) "
|
||||
"VALUES (?, ?, ?, ?)",
|
||||
[(scope, k, v, now) for k, v in entries.items() if k and v])
|
||||
|
||||
self._execute_write(_do)
|
||||
|
||||
def load_gateway_routing_entries(self, *, scope: str = "") -> Dict[str, str]:
|
||||
@@ -464,7 +460,6 @@ class SessionGatewayMixin:
|
||||
either row makes this a no-op. Non-NULL orphan columns are preserved."""
|
||||
if not orphan_id or not donor_id or orphan_id == donor_id:
|
||||
return False
|
||||
|
||||
def _do(conn):
|
||||
donor = conn.execute(
|
||||
"SELECT session_key, chat_id, chat_type, thread_id, user_id, "
|
||||
@@ -497,14 +492,12 @@ class SessionGatewayMixin:
|
||||
"end_reason = 'superseded_by_repair' WHERE id = ?",
|
||||
(time.time(), donor_id))
|
||||
return True
|
||||
|
||||
return self._execute_write(_do)
|
||||
|
||||
def increment_hygiene_failure_streak(self, session_key: str) -> int:
|
||||
"""Atomically increment the session-hygiene failure streak for one chat."""
|
||||
if not session_key:
|
||||
return 1
|
||||
|
||||
def _do(conn):
|
||||
conn.execute(
|
||||
"""INSERT INTO gateway_hygiene_state (session_key, failure_streak)
|
||||
@@ -517,7 +510,6 @@ class SessionGatewayMixin:
|
||||
"SELECT failure_streak FROM gateway_hygiene_state WHERE session_key = ?", (session_key,),
|
||||
).fetchone()
|
||||
return int(row[0])
|
||||
|
||||
return self._execute_write(_do)
|
||||
|
||||
def reset_hygiene_failure_streak(self, session_key: str) -> None:
|
||||
@@ -591,7 +583,6 @@ class SessionGatewayMixin:
|
||||
if max_age_seconds <= 0:
|
||||
return []
|
||||
cutoff = time.time() - max_age_seconds
|
||||
|
||||
def _do(conn):
|
||||
cur = conn.execute(
|
||||
"DELETE FROM gateway_heartbeats WHERE last_heartbeat < ? RETURNING backend_id",
|
||||
|
||||
@@ -293,7 +293,6 @@ class SessionMessagesMixin:
|
||||
num_tool_calls = _tool_calls_count(tool_calls)
|
||||
params = self._message_row_params(
|
||||
session_id, role, msg, tool_calls, _coerce_timestamp(timestamp, time.time()), keep_reasoning=True)
|
||||
|
||||
def _do(conn):
|
||||
self._check_transcript_write_guards(
|
||||
conn, session_id, compression_lock_holder,
|
||||
@@ -322,7 +321,6 @@ class SessionMessagesMixin:
|
||||
compression_lock_holder=compression_lock_holder, turn_lease_holder=turn_lease_holder,
|
||||
turn_lease_ttl_seconds=turn_lease_ttl_seconds)
|
||||
for start in range(0, len(messages), chunk_rows))
|
||||
|
||||
def _do(conn):
|
||||
self._check_transcript_write_guards(
|
||||
conn, session_id, compression_lock_holder,
|
||||
@@ -348,7 +346,6 @@ class SessionMessagesMixin:
|
||||
from hermes_state import _scrub_surrogates
|
||||
if not session_id or not content or not display_kind:
|
||||
return False
|
||||
|
||||
def _do(conn):
|
||||
row = conn.execute(
|
||||
"SELECT id FROM messages WHERE session_id = ? AND role = ? "
|
||||
@@ -376,7 +373,6 @@ class SessionMessagesMixin:
|
||||
from hermes_state import _scrub_surrogates
|
||||
if not session_id or message_row_id is None:
|
||||
return None
|
||||
|
||||
def _do(conn):
|
||||
row = conn.execute(_DISPLAY_META_ROW_SQL, (message_row_id, session_id)).fetchone()
|
||||
if row is None:
|
||||
@@ -409,7 +405,6 @@ class SessionMessagesMixin:
|
||||
cache-safe); the ``seen`` stamp makes each announcement exactly once."""
|
||||
if not session_id:
|
||||
return []
|
||||
|
||||
def _do(conn):
|
||||
rows = conn.execute(
|
||||
"SELECT id, role, content, display_metadata FROM messages "
|
||||
@@ -493,7 +488,6 @@ class SessionMessagesMixin:
|
||||
turn_lease`` runs the lease check in-txn for user rewrites that don't own it."""
|
||||
from hermes_state import CompressionSessionClosedError
|
||||
active_clause = " AND active = 1" if active_only else ""
|
||||
|
||||
def _do(conn):
|
||||
if reject_active_turn_lease:
|
||||
self._check_transcript_write_guards(
|
||||
@@ -563,7 +557,6 @@ class SessionMessagesMixin:
|
||||
search doesn't return each carried message once per compaction.
|
||||
``model_config_patch`` merges in the same txn (``None`` removes a key)."""
|
||||
from hermes_state import SessionCompressionInProgressError
|
||||
|
||||
def _do(conn):
|
||||
if lock_holder is not None:
|
||||
lock_row = conn.execute(_COMPRESSION_LOCK_ROW_SQL, (session_id,)).fetchone()
|
||||
@@ -1115,7 +1108,6 @@ class SessionMessagesMixin:
|
||||
head_row = conn.execute(
|
||||
"SELECT MAX(id) FROM messages WHERE session_id = ? AND active = 1", (session_id,)).fetchone()
|
||||
return target_row, ids, head_row[0] if head_row else None, replacement_message_id
|
||||
|
||||
target_row, rewound, new_head_id, replacement_message_id = self._execute_write(_do)
|
||||
# Decode for the prompt-buffer prefill without a second fallible DB operation.
|
||||
target_row["content"] = self._decode_content(target_row.get("content"))
|
||||
@@ -1192,7 +1184,6 @@ class SessionMessagesMixin:
|
||||
is touched. ``backup``: ``VACUUM INTO`` snapshot first (none when nothing changes). Returns
|
||||
``{"dry_run", "rows_affected", "row_ids", "backup_path"}``."""
|
||||
from hermes_state import _STALE_TOOL_CALL_MARKER_RE
|
||||
|
||||
def _find_affected(conn) -> List[int]:
|
||||
cursor = conn.execute(
|
||||
"SELECT id, content FROM messages "
|
||||
@@ -1200,7 +1191,6 @@ class SessionMessagesMixin:
|
||||
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}
|
||||
|
||||
@@ -1217,7 +1207,6 @@ class SessionMessagesMixin:
|
||||
self._conn.execute("VACUUM INTO ?", (str(dest),))
|
||||
backup_path = str(dest)
|
||||
logger.info("Backed up state.db to %s before clean-markers write", backup_path)
|
||||
|
||||
def _do(conn):
|
||||
ids = _find_affected(conn)
|
||||
if ids:
|
||||
|
||||
Reference in New Issue
Block a user