From ffef90ae7de2248d659252b408d8fc32e2b5098c Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:00:07 -0700 Subject: [PATCH] =?UTF-8?q?refactor(state):=20gateway=20mixin=20=E2=80=94?= =?UTF-8?q?=20hoisted=20peer/orphan=20SQL=20constants,=20dead=20tuple-row?= =?UTF-8?q?=20branches=20removed,=20=5Fwrite=5Frowcount=20for=20fail=5Fhan?= =?UTF-8?q?doff;=20maintenance=20table-driven=20prune=20filters?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- hermes_state_gateway.py | 611 ++++++++++++++---------------------- hermes_state_maintenance.py | 440 +++++++++----------------- 2 files changed, 378 insertions(+), 673 deletions(-) diff --git a/hermes_state_gateway.py b/hermes_state_gateway.py index f74b7dcce0..bbf0f09b81 100644 --- a/hermes_state_gateway.py +++ b/hermes_state_gateway.py @@ -8,7 +8,6 @@ from __future__ import annotations import json import logging -import sqlite3 import sys import time from pathlib import Path @@ -23,6 +22,125 @@ from hermes_state_common import ( # Log-record parity with the origin module (caplog tests pin "hermes_state"). logger = logging.getLogger("hermes_state") +# Recursive CTE naming a session plus its compression ancestors (rows a +# resume must keep on one routing peer); branch/delegate/tool rows stop it. +_COMPRESSION_LINEAGE_CTE = """ + WITH RECURSIVE compression_lineage(id) AS ( + SELECT ? + UNION + SELECT parent.id + FROM compression_lineage lineage + JOIN sessions child ON child.id = lineage.id + JOIN sessions parent ON parent.id = child.parent_session_id + WHERE parent.end_reason = 'compression' + AND json_extract( + COALESCE(child.model_config, '{}'), + '$._branched_from' + ) IS NULL + AND json_extract( + COALESCE(child.model_config, '{}'), + '$._delegate_from' + ) IS NULL + AND COALESCE(child.source, '') != 'tool' + ) + """ + +# Projection shared by both peer-recovery queries (exact key, then peer tuple). +_PEER_SELECT_HEAD = """ + SELECT s.*, + COALESCE(sp.prompt, s.system_prompt) + AS _system_prompt_resolved, + (COALESCE(s.message_count, 0) > 0 OR EXISTS ( + SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1 + )) AS _has_messages + FROM sessions s + LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash +""" +_PEER_BY_KEY_SQL = f"""{_PEER_SELECT_HEAD} WHERE s.session_key = ? + AND s.source = ? + AND (s.ended_at IS NULL OR s.end_reason IN ({_RECOVERABLE_END_REASONS_SQL})) + AND NOT EXISTS ( + SELECT 1 FROM sessions b + WHERE b.session_key = s.session_key + AND b.source = s.source + AND b.ended_at IS NOT NULL + AND b.end_reason IN ({_RESET_END_REASONS_SQL}) + AND b.ended_at + > COALESCE(s.last_activity_at, s.started_at) + ) + ORDER BY _has_messages DESC, + COALESCE(s.last_activity_at, s.started_at) DESC + LIMIT 1 + """ +_PEER_BY_TUPLE_SQL = f"""{_PEER_SELECT_HEAD} WHERE s.source = ? + AND COALESCE(s.user_id, '') = COALESCE(?, '') + AND COALESCE(s.chat_id, '') = COALESCE(?, '') + AND COALESCE(s.chat_type, '') = COALESCE(?, '') + AND COALESCE(s.thread_id, '') = COALESCE(?, '') + AND (? IS NULL OR COALESCE(s.profile_name, ?) = ?) + AND (s.ended_at IS NULL OR s.end_reason IN ({_RECOVERABLE_END_REASONS_SQL})) + AND (COALESCE(s.message_count, 0) > 0 OR EXISTS ( + SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1 + )) + AND NOT EXISTS ( + SELECT 1 FROM sessions b + WHERE b.source = s.source + AND COALESCE(b.user_id, '') = COALESCE(s.user_id, '') + AND COALESCE(b.chat_id, '') = COALESCE(s.chat_id, '') + AND COALESCE(b.chat_type, '') = COALESCE(s.chat_type, '') + AND COALESCE(b.thread_id, '') = COALESCE(s.thread_id, '') + AND b.ended_at IS NOT NULL + AND b.end_reason IN ({_RESET_END_REASONS_SQL}) + AND b.ended_at + > COALESCE(s.last_activity_at, s.started_at) + ) + ORDER BY COALESCE(s.last_activity_at, s.started_at) DESC + LIMIT 1 + """ + +_ORPHAN_DONOR_COLUMNS = ( + "d.id, d.session_key, d.chat_id, d.chat_type, d.thread_id, " + "d.user_id, d.origin_json, d.display_name, d.end_reason" +) +_ORPHANS_SQL = f""" + SELECT o.id, o.source, o.user_id, o.started_at, + o.parent_session_id, + {_sql_session_last_active("o")} AS last_active, + (SELECT COUNT(*) FROM messages m + WHERE m.session_id = o.id) AS message_count + FROM sessions o + WHERE o.session_key IS NULL + AND EXISTS (SELECT 1 FROM messages m + WHERE m.session_id = o.id) + AND COALESCE(o.source, '') != 'tool' + AND json_extract(COALESCE(o.model_config, '{{}}'), + '$._branched_from') IS NULL + AND json_extract(COALESCE(o.model_config, '{{}}'), + '$._delegate_from') IS NULL + ORDER BY o.started_at ASC + """ +_ORPHAN_LINEAGE_DONOR_SQL = f""" + SELECT {_ORPHAN_DONOR_COLUMNS} + FROM sessions d + WHERE d.id = ? + AND d.session_key IS NOT NULL + AND COALESCE(d.source, '') = COALESCE(?, '') + """ +_ORPHAN_CONTIGUITY_DONORS_SQL = f""" + SELECT {_ORPHAN_DONOR_COLUMNS}, {_sql_session_last_active("d")} AS last_active + FROM sessions d + WHERE d.session_key IS NOT NULL + AND d.id != ? + AND COALESCE(d.source, '') = COALESCE(?, '') + AND (COALESCE(d.user_id, '') = '' + OR COALESCE(?, '') = '' + OR d.user_id = ?) + AND {_sql_session_last_active("d")} BETWEEN ? AND ? + AND {_sql_session_last_active("d")} < ? + ORDER BY last_active DESC + LIMIT 2 + """ + class SessionGatewayMixin: """Routing index, session peers/orphans, hygiene streaks, heartbeats, handoffs.""" @@ -48,12 +166,9 @@ class SessionGatewayMixin: for pid in _concrete_state_db_holder_pids(self.db_path, holders): try: process = psutil.Process(pid) - statuses = [ - conn.status for conn in process.net_connections(kind="inet") - ] + statuses = [conn.status for conn in process.net_connections(kind="inet")] if not _is_inactive_orphan_desktop_holder( - ppid=process.ppid(), - age_seconds=now - process.create_time(), + ppid=process.ppid(), age_seconds=now - process.create_time(), min_age_seconds=min_age_seconds, ephemeral_backend=_is_ephemeral_port_zero_backend(process.cmdline()), connection_statuses=statuses, @@ -72,7 +187,6 @@ class SessionGatewayMixin: continue if not signalled: return [] - try: _gone, alive = psutil.wait_procs(candidates, timeout=1.5) except Exception: @@ -90,17 +204,9 @@ class SessionGatewayMixin: return signalled def record_gateway_session_peer( - self, - session_id: str, - *, - source: str, - user_id: str = None, - session_key: str = None, - chat_id: str = None, - chat_type: str = None, - thread_id: str = None, - display_name: str = None, - origin_json: str = None, + self, session_id: str, *, source: str, user_id: str = None, session_key: str = None, + chat_id: str = None, chat_type: str = None, thread_id: str = None, + display_name: str = None, origin_json: str = None, include_compression_ancestors: bool = False, ) -> None: """Persist the gateway routing peer for an existing session row. @@ -120,48 +226,17 @@ class SessionGatewayMixin: """ if not session_id or not session_key: return - - def _do(conn): + identity = (session_key, source, user_id, chat_id, chat_type, thread_id, display_name, origin_json) + if include_compression_ancestors: + lineage_cte = _COMPRESSION_LINEAGE_CTE + target_clause = "WHERE id IN (SELECT id FROM compression_lineage)" + query_params = [session_id, *identity] + else: lineage_cte = "" target_clause = "WHERE id = ?" - query_params = [] - if include_compression_ancestors: - lineage_cte = """ - WITH RECURSIVE compression_lineage(id) AS ( - SELECT ? - UNION - SELECT parent.id - FROM compression_lineage lineage - JOIN sessions child ON child.id = lineage.id - JOIN sessions parent ON parent.id = child.parent_session_id - WHERE parent.end_reason = 'compression' - AND json_extract( - COALESCE(child.model_config, '{}'), - '$._branched_from' - ) IS NULL - AND json_extract( - COALESCE(child.model_config, '{}'), - '$._delegate_from' - ) IS NULL - AND COALESCE(child.source, '') != 'tool' - ) - """ - target_clause = "WHERE id IN (SELECT id FROM compression_lineage)" - query_params.append(session_id) - query_params.extend( - ( - session_key, - source, - user_id, - chat_id, - chat_type, - thread_id, - display_name, - origin_json, - ) - ) - if not include_compression_ancestors: - query_params.append(session_id) + query_params = [*identity, session_id] + + def _do(conn): conn.execute( f"""{lineage_cte} UPDATE sessions @@ -174,13 +249,12 @@ class SessionGatewayMixin: ) # Self-heal: the UPDATE silently no-ops on a missing row — insert it # with full identity so the session is durably routable. - if not include_compression_ancestors: - cur = conn.execute( - "SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,) - ) - if cur.fetchone() is None: - conn.execute( - """INSERT INTO sessions ( + if include_compression_ancestors: + return + cur = conn.execute("SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,)) + if cur.fetchone() is None: + conn.execute( + """INSERT INTO sessions ( id, source, user_id, session_key, chat_id, chat_type, thread_id, display_name, origin_json, profile_name, started_at @@ -193,28 +267,19 @@ class SessionGatewayMixin: thread_id = COALESCE(sessions.thread_id, excluded.thread_id), display_name = COALESCE(sessions.display_name, excluded.display_name), origin_json = COALESCE(sessions.origin_json, excluded.origin_json)""", - ( - session_id, - source, - user_id, - session_key, - chat_id, - chat_type, - thread_id, - display_name, - origin_json, - # Same ownership stamp as _insert_session_row: an - # unowned (NULL) row vanishes from profile-keyed consumers. - self._own_profile_name(), - time.time(), - ), - ) + ( + session_id, source, user_id, session_key, chat_id, chat_type, thread_id, + display_name, origin_json, + # Same ownership stamp as _insert_session_row: an + # unowned (NULL) row vanishes from profile-keyed consumers. + 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: + def save_gateway_routing_entry(self, session_key: str, entry_json: str, *, scope: str = "") -> None: """Upsert one gateway routing entry (session_key -> SessionEntry JSON). ``gateway_routing`` durably replaces sessions.json. ``scope`` namespaces @@ -222,7 +287,6 @@ class SessionGatewayMixin: """ if not session_key or not entry_json: return - self._write_sql( """INSERT INTO gateway_routing (scope, session_key, entry_json, updated_at) VALUES (?, ?, ?, ?) @@ -232,9 +296,7 @@ class SessionGatewayMixin: (scope, session_key, entry_json, time.time()), ) - def replace_gateway_routing_entries( - self, entries: Dict[str, str], *, scope: str = "" - ) -> None: + def replace_gateway_routing_entries(self, entries: Dict[str, str], *, scope: str = "") -> None: """Atomically replace the routing index for *scope* with *entries*. Full-rewrite semantics: keys absent from *entries* are removed. One @@ -256,14 +318,11 @@ class SessionGatewayMixin: def load_gateway_routing_entries(self, *, scope: str = "") -> Dict[str, str]: """Load routing entries for *scope* as {session_key: entry_json}.""" rows = self._read_all( - "SELECT session_key, entry_json FROM gateway_routing WHERE scope = ?", - (scope,), + "SELECT session_key, entry_json FROM gateway_routing WHERE scope = ?", (scope,) ) return {r["session_key"]: r["entry_json"] for r in rows} - def list_never_active_keyed_sessions( - self, *, older_than_days: float - ) -> List[Dict[str, Any]]: + def list_never_active_keyed_sessions(self, *, older_than_days: float) -> List[Dict[str, Any]]: """Keyed gateway rows that were opened and then never used at all. Keyed, still-open rows with no evidence of a single turn (no messages, @@ -321,19 +380,13 @@ class SessionGatewayMixin: doomed.append((row["scope"], row["session_key"])) if not doomed: return 0 - self._write_sql( - "DELETE FROM gateway_routing WHERE scope = ? AND session_key = ?", - doomed, - many=True, + "DELETE FROM gateway_routing WHERE scope = ? AND session_key = ?", doomed, many=True ) return len(doomed) def prune_never_active_keyed_sessions( - self, - *, - older_than_days: float, - sessions_dir: Optional[Path] = None, + self, *, older_than_days: float, sessions_dir: Optional[Path] = None ) -> Tuple[int, int]: """Delete never-active keyed rows and the routing entries naming them. @@ -343,24 +396,16 @@ class SessionGatewayMixin: :meth:`delete_session` so the delegate cascade, FTS bookkeeping and transcript cleanup stay owned by one implementation. """ - candidates = self.list_never_active_keyed_sessions( - older_than_days=older_than_days - ) + candidates = self.list_never_active_keyed_sessions(older_than_days=older_than_days) if not candidates: return (0, 0) ids = {str(row["id"]) for row in candidates} routing_deleted = self._delete_routing_entries_for_sessions(ids) - deleted = 0 - for session_id in ids: - if self.delete_session(session_id, sessions_dir=sessions_dir): - deleted += 1 + deleted = sum(1 for sid in ids if self.delete_session(sid, sessions_dir=sessions_dir)) return (deleted, routing_deleted) def list_gateway_sessions( - self, - *, - platform: Optional[str] = None, - active_only: bool = True, + self, *, platform: Optional[str] = None, active_only: bool = True ) -> List[Dict[str, Any]]: """List gateway sessions (rows with a session_key): newest row per key, one live mapping per routing key. ``platform`` filters on ``source``.""" @@ -388,17 +433,11 @@ class SessionGatewayMixin: if active_only: query += " AND ended_at IS NULL" query += " ORDER BY last_active DESC" - rows = self._read_all(query, params) - return [self._session_row_dict(r) for r in rows] + return [self._session_row_dict(r) for r in self._read_all(query, params)] def find_latest_gateway_session_for_peer( - self, - *, - source: str, - user_id: Optional[str] = None, - session_key: Optional[str] = None, - chat_id: Optional[str] = None, - chat_type: Optional[str] = None, + self, *, source: str, user_id: Optional[str] = None, session_key: Optional[str] = None, + chat_id: Optional[str] = None, chat_type: Optional[str] = None, thread_id: Optional[str] = None, ) -> Optional[Dict[str, Any]]: """Find the latest recoverable gateway session for a routing peer. @@ -421,93 +460,30 @@ class SessionGatewayMixin: the same peer, or the has-messages ranking could reach behind a /new and restore the exact context the user reset — so a candidate is rejected when a peer boundary row ended *after* its last activity. + + Fallback for a temporarily-missing exact key still requires the + complete peer tuple (never cross chats/threads/users) and a profile + fence: a Telegram DM's peer tuple is identical for every bot (chat_id + == user_id, no thread), so a sibling profile's legacy row would + otherwise be adopted. A row is ours when profile_name is the owner or + NULL; stores outside the profile tree derive no owner and stay unfenced. """ if not session_key: return None with self._read_ctx() as conn: - row = conn.execute( - f""" - SELECT s.*, - COALESCE(sp.prompt, s.system_prompt) - AS _system_prompt_resolved, - (COALESCE(s.message_count, 0) > 0 OR EXISTS ( - SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1 - )) AS _has_messages - FROM sessions s - LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash - WHERE s.session_key = ? - AND s.source = ? - AND (s.ended_at IS NULL OR s.end_reason IN ({_RECOVERABLE_END_REASONS_SQL})) - AND NOT EXISTS ( - SELECT 1 FROM sessions b - WHERE b.session_key = s.session_key - AND b.source = s.source - AND b.ended_at IS NOT NULL - AND b.end_reason IN ({_RESET_END_REASONS_SQL}) - AND b.ended_at - > COALESCE(s.last_activity_at, s.started_at) - ) - ORDER BY _has_messages DESC, - COALESCE(s.last_activity_at, s.started_at) DESC - LIMIT 1 - """, - (session_key, source), - ).fetchone() + row = conn.execute(_PEER_BY_KEY_SQL, (session_key, source)).fetchone() if row is not None: return self._session_row_dict(row) - - # Conservative fallback for a temporarily-missing exact key: still - # require the complete peer tuple so we never cross chats/threads/users. if chat_id is None or chat_type is None: return None - # Profile fence: a Telegram DM's peer tuple is identical for every - # bot (chat_id == user_id, no thread), so a sibling profile's legacy - # row would otherwise be adopted. A row is ours when profile_name is - # the owner or NULL (legacy rows this store minted); stores outside - # the profile tree derive no owner and stay unfenced. owner = self._own_profile_name() row = conn.execute( - f""" - SELECT s.*, - COALESCE(sp.prompt, s.system_prompt) - AS _system_prompt_resolved, - (COALESCE(s.message_count, 0) > 0 OR EXISTS ( - SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1 - )) AS _has_messages - FROM sessions s - LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash - WHERE s.source = ? - AND COALESCE(s.user_id, '') = COALESCE(?, '') - AND COALESCE(s.chat_id, '') = COALESCE(?, '') - AND COALESCE(s.chat_type, '') = COALESCE(?, '') - AND COALESCE(s.thread_id, '') = COALESCE(?, '') - AND (? IS NULL OR COALESCE(s.profile_name, ?) = ?) - AND (s.ended_at IS NULL OR s.end_reason IN ({_RECOVERABLE_END_REASONS_SQL})) - AND (COALESCE(s.message_count, 0) > 0 OR EXISTS ( - SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1 - )) - AND NOT EXISTS ( - SELECT 1 FROM sessions b - WHERE b.source = s.source - AND COALESCE(b.user_id, '') = COALESCE(s.user_id, '') - AND COALESCE(b.chat_id, '') = COALESCE(s.chat_id, '') - AND COALESCE(b.chat_type, '') = COALESCE(s.chat_type, '') - AND COALESCE(b.thread_id, '') = COALESCE(s.thread_id, '') - AND b.ended_at IS NOT NULL - AND b.end_reason IN ({_RESET_END_REASONS_SQL}) - AND b.ended_at - > COALESCE(s.last_activity_at, s.started_at) - ) - ORDER BY COALESCE(s.last_activity_at, s.started_at) DESC - LIMIT 1 - """, + _PEER_BY_TUPLE_SQL, (source, user_id, chat_id, chat_type, thread_id, owner, owner, owner), ).fetchone() return self._session_row_dict(row) if row else None - def find_orphaned_gateway_sessions( - self, *, max_gap_s: Optional[float] = None - ) -> List[Dict[str, Any]]: + def find_orphaned_gateway_sessions(self, *, max_gap_s: Optional[float] = None) -> List[Dict[str, Any]]: """Report message-bearing session rows that lost their routing identity. A candidate orphan has messages but no ``session_key``; it is @@ -523,137 +499,56 @@ class SessionGatewayMixin: mis-adopting splices one person's conversation into another's chat. Branch/delegate/tool rows are excluded — unkeyed by design, not damage. """ - gap = ( - self._ORPHAN_ADOPTION_MAX_GAP_S - if max_gap_s is None - else float(max_gap_s) - ) - orphan_active = _sql_session_last_active("o") - donor_active = _sql_session_last_active("d") - donor_columns = ( - "d.id, d.session_key, d.chat_id, d.chat_type, d.thread_id, " - "d.user_id, d.origin_json, d.display_name, d.end_reason" - ) + gap = self._ORPHAN_ADOPTION_MAX_GAP_S if max_gap_s is None else float(max_gap_s) records: List[Dict[str, Any]] = [] - with self._read_ctx() as conn: - orphans = conn.execute( - f""" - SELECT o.id, o.source, o.user_id, o.started_at, - o.parent_session_id, - {orphan_active} AS last_active, - (SELECT COUNT(*) FROM messages m - WHERE m.session_id = o.id) AS message_count - FROM sessions o - WHERE o.session_key IS NULL - AND EXISTS (SELECT 1 FROM messages m - WHERE m.session_id = o.id) - AND COALESCE(o.source, '') != 'tool' - AND json_extract(COALESCE(o.model_config, '{{}}'), - '$._branched_from') IS NULL - AND json_extract(COALESCE(o.model_config, '{{}}'), - '$._delegate_from') IS NULL - ORDER BY o.started_at ASC - """ - ).fetchall() - - for orphan in orphans: + for orphan in conn.execute(_ORPHANS_SQL).fetchall(): donor = None - evidence = "" reason = "" - if orphan["parent_session_id"]: evidence = "lineage" donor = conn.execute( - f""" - SELECT {donor_columns} - FROM sessions d - WHERE d.id = ? - AND d.session_key IS NOT NULL - AND COALESCE(d.source, '') = COALESCE(?, '') - """, - (orphan["parent_session_id"], orphan["source"]), + _ORPHAN_LINEAGE_DONOR_SQL, (orphan["parent_session_id"], orphan["source"]) ).fetchone() if donor is None: - reason = ( - "parent session carries no gateway identity of " - "this source" - ) + reason = "parent session carries no gateway identity of this source" else: evidence = "contiguity" + started = orphan["started_at"] or 0 candidates = conn.execute( - f""" - SELECT {donor_columns}, {donor_active} AS last_active - FROM sessions d - WHERE d.session_key IS NOT NULL - AND d.id != ? - AND COALESCE(d.source, '') = COALESCE(?, '') - AND (COALESCE(d.user_id, '') = '' - OR COALESCE(?, '') = '' - OR d.user_id = ?) - AND {donor_active} BETWEEN ? AND ? - AND {donor_active} < ? - ORDER BY last_active DESC - LIMIT 2 - """, - ( - orphan["id"], - orphan["source"], - orphan["user_id"], - orphan["user_id"], - (orphan["started_at"] or 0) - gap, - (orphan["started_at"] or 0) + gap, - orphan["last_active"], - ), + _ORPHAN_CONTIGUITY_DONORS_SQL, + (orphan["id"], orphan["source"], orphan["user_id"], orphan["user_id"], + started - gap, started + gap, orphan["last_active"]), ).fetchall() if not candidates: - reason = ( - f"no keyed predecessor fell quiet within {gap:.0f}s " - "of this session's start" - ) + reason = f"no keyed predecessor fell quiet within {gap:.0f}s of this session's start" elif len(candidates) > 1: - reason = ( - "ambiguous: more than one keyed predecessor " - "matches this window" - ) + reason = "ambiguous: more than one keyed predecessor matches this window" else: donor = candidates[0] - - records.append( - { - "orphan_id": orphan["id"], - "source": orphan["source"], - "message_count": orphan["message_count"], - "started_at": orphan["started_at"], - "last_active": orphan["last_active"], - "donor_id": donor["id"] if donor else None, - "session_key": donor["session_key"] if donor else None, - "evidence": evidence if donor else "", - "adoptable": donor is not None, - "reason": reason, - } - ) - + records.append({ + "orphan_id": orphan["id"], "source": orphan["source"], + "message_count": orphan["message_count"], "started_at": orphan["started_at"], + "last_active": orphan["last_active"], + "donor_id": donor["id"] if donor else None, + "session_key": donor["session_key"] if donor else None, + "evidence": evidence if donor else "", + "adoptable": donor is not None, + "reason": reason, + }) # Two unkeyed successors claiming one predecessor: at most one continues # that chat, and nothing here says which. contested = { - r["donor_id"] - for r in records - if r["adoptable"] - and sum(1 for x in records if x["donor_id"] == r["donor_id"]) > 1 + r["donor_id"] for r in records + if r["adoptable"] and sum(1 for x in records if x["donor_id"] == r["donor_id"]) > 1 } for record in records: if record["donor_id"] in contested: record["adoptable"] = False - record["reason"] = ( - "ambiguous: more than one unkeyed session claims this " - "predecessor" - ) + record["reason"] = "ambiguous: more than one unkeyed session claims this predecessor" return records - def adopt_orphaned_gateway_session( - self, orphan_id: str, donor_id: str - ) -> bool: + def adopt_orphaned_gateway_session(self, orphan_id: str, donor_id: str) -> bool: """Stamp *orphan_id* with *donor_id*'s routing identity, retire *donor_id*. Re-verifies the pair inside the write transaction so a concurrent @@ -670,8 +565,7 @@ class SessionGatewayMixin: (donor_id,), ).fetchone() orphan = conn.execute( - "SELECT session_key, source FROM sessions WHERE id = ?", - (orphan_id,), + "SELECT session_key, source FROM sessions WHERE id = ?", (orphan_id,) ).fetchone() if donor is None or orphan is None: return False @@ -679,7 +573,6 @@ class SessionGatewayMixin: return False if (donor["source"] or "") != (orphan["source"] or ""): return False - conn.execute( """UPDATE sessions SET session_key = ?, @@ -691,17 +584,8 @@ class SessionGatewayMixin: display_name = COALESCE(display_name, ?), parent_session_id = COALESCE(parent_session_id, ?) WHERE id = ? AND session_key IS NULL""", - ( - donor["session_key"], - donor["chat_id"], - donor["chat_type"], - donor["thread_id"], - donor["user_id"], - donor["origin_json"], - donor["display_name"], - donor_id, - orphan_id, - ), + (donor["session_key"], donor["chat_id"], donor["chat_type"], donor["thread_id"], + donor["user_id"], donor["origin_json"], donor["display_name"], donor_id, orphan_id), ) # Retire the predecessor under a reason recovery does NOT treat as # resumable — 'agent_close'/'ws_orphan_reap' would keep it in the @@ -719,7 +603,6 @@ class SessionGatewayMixin: """Atomically increment the session-hygiene failure streak for one chat.""" if not session_key: return 1 - result = [] def _do(conn): conn.execute( @@ -733,20 +616,15 @@ class SessionGatewayMixin: "SELECT failure_streak FROM gateway_hygiene_state WHERE session_key = ?", (session_key,), ).fetchone() - result.append(int(row[0])) + return int(row[0]) - self._execute_write(_do) - return result[0] + return self._execute_write(_do) def reset_hygiene_failure_streak(self, session_key: str) -> None: """Clear the persisted session-hygiene failure streak for one chat.""" if not session_key: return - - self._write_sql( - "DELETE FROM gateway_hygiene_state WHERE session_key = ?", - (session_key,), - ) + self._write_sql("DELETE FROM gateway_hygiene_state WHERE session_key = ?", (session_key,)) @staticmethod def session_gateway_runtime(session_meta: Optional[Dict[str, Any]]) -> Dict[str, Any]: @@ -769,41 +647,26 @@ class SessionGatewayMixin: if not isinstance(raw, dict): raw = {} runtime = raw.get("gateway_runtime") + # Filter None: the persist path writes or-None to trigger deletion in + # the top-level merge, but gateway_runtime is replaced whole (not + # deep-merged), so None values survive here. if isinstance(runtime, dict) and runtime.get("provider"): - # Filter None: the persist path writes or-None to trigger deletion - # in the top-level merge, but gateway_runtime is replaced whole - # (not deep-merged), so None values survive here. return {k: v for k, v in runtime.items() if v is not None} - top_level = { - key: raw.get(key) - for key in ("provider", "base_url", "api_mode") - if raw.get(key) - } + top_level = {key: raw.get(key) for key in ("provider", "base_url", "api_mode") if raw.get(key)} if top_level: return top_level # Last resort: billing_provider, COALESCE-written on the first accounted # API call — the only durable record for sessions that never ran /model. # Bare buckets ("auto"/"custom") are not routable identities; filter # them so resume falls back to the ambient config default. - billing_provider = str( - (session_meta or {}).get("billing_provider") or "" - ).strip() - if ( - billing_provider - and billing_provider.lower() not in _BARE_BILLING_PROVIDERS - ): + billing_provider = str((session_meta or {}).get("billing_provider") or "").strip() + if billing_provider and billing_provider.lower() not in _BARE_BILLING_PROVIDERS: return {"provider": billing_provider} return {k: v for k, v in (runtime or {}).items() if v is not None} if isinstance(runtime, dict) else {} def register_backend_heartbeat( - self, - *, - backend_id: str, - pid: int, - started_at: float, - last_heartbeat: Optional[float] = None, - profile: str = "", - host: str = "", + self, *, backend_id: str, pid: int, started_at: float, + last_heartbeat: Optional[float] = None, profile: str = "", host: str = "", ) -> None: """Upsert this backend's liveness row. @@ -826,8 +689,7 @@ class SessionGatewayMixin: " last_heartbeat = excluded.last_heartbeat," " profile = excluded.profile," " host = excluded.host", - (str(backend_id), int(pid), float(started_at), ts, - str(profile), str(host)), + (str(backend_id), int(pid), float(started_at), ts, str(profile), str(host)), ) def clear_backend_heartbeat(self, backend_id: str) -> bool: @@ -836,8 +698,7 @@ class SessionGatewayMixin: if not backend_id: return False return self._write_rowcount( - "DELETE FROM gateway_heartbeats WHERE backend_id = ?", - (str(backend_id),), + "DELETE FROM gateway_heartbeats WHERE backend_id = ?", (str(backend_id),) ) > 0 def prune_stale_heartbeats(self, *, max_age_seconds: float) -> List[str]: @@ -846,6 +707,7 @@ 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 < ?" @@ -862,16 +724,7 @@ class SessionGatewayMixin: " profile, host FROM gateway_heartbeats" " ORDER BY last_heartbeat DESC", ) - out: List[Dict[str, Any]] = [] - for r in rows: - if isinstance(r, sqlite3.Row): - out.append({k: r[k] for k in r.keys()}) - else: - out.append({ - "backend_id": r[0], "pid": r[1], "started_at": r[2], - "last_heartbeat": r[3], "profile": r[4], "host": r[5], - }) - return out + return [dict(r) for r in rows] def request_handoff(self, session_id: str, platform: str) -> bool: """Mark a session pending handoff to *platform*; False if a handoff is already in flight.""" @@ -935,11 +788,7 @@ class SessionGatewayMixin: ) def fail_handoff( - self, - session_id: str, - error: str, - *, - only_states: Optional[Tuple[str, ...]] = None, + self, session_id: str, error: str, *, only_states: Optional[Tuple[str, ...]] = None ) -> bool: """Mark a handoff failed and record the reason; True when a row transitioned. @@ -952,22 +801,20 @@ class SessionGatewayMixin: (split-brain: the handoff delivered and ``switch_session`` re-pointed the session). The watcher fails its OWN claimed row unconditionally. """ - def _do(conn): - if only_states: - placeholders = ", ".join("?" for _ in only_states) - cur = conn.execute( - "UPDATE sessions SET handoff_state = 'failed', " - f"handoff_error = ? WHERE id = ? AND handoff_state IN ({placeholders})", - (error[:500], session_id, *only_states), - ) - else: - cur = conn.execute( - "UPDATE sessions SET handoff_state = 'failed', " - "handoff_error = ? WHERE id = ?", - (error[:500], session_id), - ) - return cur.rowcount > 0 - return bool(self._execute_write(_do)) + if only_states: + placeholders = ", ".join("?" for _ in only_states) + sql = ( + "UPDATE sessions SET handoff_state = 'failed', " + f"handoff_error = ? WHERE id = ? AND handoff_state IN ({placeholders})" + ) + params = (error[:500], session_id, *only_states) + else: + sql = ( + "UPDATE sessions SET handoff_state = 'failed', " + "handoff_error = ? WHERE id = ?" + ) + params = (error[:500], session_id) + return self._write_rowcount(sql, params) > 0 def reclaim_stale_running_handoffs(self, error: str) -> List[str]: """Fail every handoff stuck in ``running``. Returns the ids reclaimed. @@ -982,9 +829,7 @@ class SessionGatewayMixin: delivery; a clean terminal state the user can retry from is right. """ def _do(conn): - cur = conn.execute( - "SELECT id FROM sessions WHERE handoff_state = 'running'" - ) + cur = conn.execute("SELECT id FROM sessions WHERE handoff_state = 'running'") ids = [r[0] for r in cur.fetchall()] if ids: conn.execute( diff --git a/hermes_state_maintenance.py b/hermes_state_maintenance.py index 89829b9191..08ef9db6ab 100644 --- a/hermes_state_maintenance.py +++ b/hermes_state_maintenance.py @@ -17,6 +17,63 @@ from hermes_state_common import ( # caplog tests pin the "hermes_state" logger name. logger = logging.getLogger("hermes_state") +_LAST_ACTIVE_SQL = """COALESCE( + (SELECT MAX(m.timestamp) FROM messages m + WHERE m.session_id = s.id), + s.started_at + )""" +_TOKENS_SQL = "(COALESCE(s.input_tokens, 0) + COALESCE(s.output_tokens, 0))" +_COST_SQL = "COALESCE(s.actual_cost_usd, s.estimated_cost_usd, 0)" + + +def _like(value: str) -> str: + return f"%{_escape_like(value.lower())}%" + + +def _cwd_prefix_filter(value: str) -> Tuple[List[str], list]: + from hermes_state import _cwd_prefix_clause + clause, params = _cwd_prefix_clause(value) + return [clause], list(params) + + +def _one(clause: str, conv=None): + return lambda v: ([clause], [conv(v) if conv else v]) + + +# Prune/archive filters in evaluation order: (kwarg, applies-when, builder). +# ``applies-when`` is "notnone" (numeric/time bounds; 0 is a real bound) or +# "truthy" (strings; "" means unset). Builders return (clauses, params). +_PRUNE_FILTERS = ( + # Orphan-swept rows age from the sweep, not their old activity, or the + # next prune pass deletes them before the user can recover. + ("last_active_before", "notnone", lambda v: ( + [_LAST_ACTIVE_SQL + " < ?", + "(COALESCE(s.end_reason, '') != 'startup_orphan_reap' OR s.ended_at < ?)"], + [v, v])), + ("last_active_after", "notnone", _one(_LAST_ACTIVE_SQL + " >= ?")), + ("started_before", "notnone", _one("s.started_at < ?")), + ("started_after", "notnone", _one("s.started_at >= ?")), + ("source", "truthy", _one("s.source = ?")), + ("title_like", "truthy", _one("LOWER(COALESCE(s.title, '')) LIKE ? ESCAPE '\\'", _like)), + ("end_reason", "truthy", _one("s.end_reason = ?")), + ("cwd_prefix", "truthy", _cwd_prefix_filter), + ("min_messages", "notnone", _one("s.message_count >= ?")), + ("max_messages", "notnone", _one("s.message_count <= ?")), + ("model_like", "truthy", _one("LOWER(COALESCE(s.model, '')) LIKE ? ESCAPE '\\'", _like)), + ("provider", "truthy", _one("LOWER(COALESCE(s.billing_provider, '')) = ?", str.lower)), + ("user_id", "truthy", _one("s.user_id = ?")), + ("chat_id", "truthy", _one("s.chat_id = ?")), + ("chat_type", "truthy", _one("s.chat_type = ?")), + ("branch_like", "truthy", _one("LOWER(COALESCE(s.git_branch, '')) LIKE ? ESCAPE '\\'", _like)), + ("min_tokens", "notnone", _one(_TOKENS_SQL + " >= ?")), + ("max_tokens", "notnone", _one(_TOKENS_SQL + " <= ?")), + ("min_cost", "notnone", _one(_COST_SQL + " >= ?")), + ("max_cost", "notnone", _one(_COST_SQL + " <= ?")), + ("min_tool_calls", "notnone", _one("COALESCE(s.tool_call_count, 0) >= ?")), + ("max_tool_calls", "notnone", _one("COALESCE(s.tool_call_count, 0) <= ?")), +) +_PRUNE_FILTER_NAMES = frozenset(name for name, _, _ in _PRUNE_FILTERS) | {"archived", "include_pinned"} + class SessionMaintenanceMixin: """Retention pruning, stale-session archiving and VACUUM policy for SessionDB.""" @@ -39,9 +96,7 @@ class SessionMaintenanceMixin: ids = [r[0] for r in rows] if ids: placeholders = ",".join("?" * len(ids)) - conn.execute( - f"DELETE FROM sessions WHERE id IN ({placeholders})", ids - ) + conn.execute(f"DELETE FROM sessions WHERE id IN ({placeholders})", ids) self._delete_unreferenced_system_prompts(conn) return ids @@ -52,12 +107,9 @@ class SessionMaintenanceMixin: return len(removed_ids) def sweep_orphaned_sessions( - self, - *, - max_idle_seconds: float, + self, *, max_idle_seconds: float, sources: Tuple[str, ...] = ("tui", "desktop", "subagent"), - exclude_ids: Tuple[str, ...] = (), - exclude_pinned: bool = False, + exclude_ids: Tuple[str, ...] = (), exclude_pinned: bool = False, heartbeat_staleness_seconds: Optional[float] = None, heartbeat_ownership_grace_seconds: Optional[float] = None, respect_gateway_heartbeats: bool = True, @@ -67,11 +119,10 @@ class SessionMaintenanceMixin: The TUI/desktop gateway reaps disconnected sessions with an in-process grace timer; a restart destroys the timer and leaves ``ended_at IS NULL`` forever. This closes rows for ``sources`` whose ``started_at`` - AND canonical last activity (newest of ``last_activity_at`` and the - newest message, else ``started_at``) are both older than - ``max_idle_seconds``, with ``end_reason='startup_orphan_reap'``. The - separate ``started_at`` predicate protects fresh compression/branch - children whose copied activity is old. + AND canonical last activity are both older than ``max_idle_seconds``, + with ``end_reason='startup_orphan_reap'``. The separate ``started_at`` + predicate protects fresh compression/branch children whose copied + activity is old. Only pass sources whose lifecycle the caller owns — never messaging platforms like ``telegram`` (ending those triggers a routing loop). @@ -103,20 +154,15 @@ class SessionMaintenanceMixin: ) hb_grace = ( heartbeat_ownership_grace_seconds - if heartbeat_ownership_grace_seconds is not None - and heartbeat_ownership_grace_seconds >= 0 + if heartbeat_ownership_grace_seconds is not None and heartbeat_ownership_grace_seconds >= 0 else hb_staleness ) now = time.time() cutoff = now - max_idle_seconds - hb_cutoff = now - hb_staleness placeholders = ",".join("?" for _ in srcs) - staleness = ( - f"started_at < ? AND {_sql_session_last_active('sessions')} < ?" - ) pin_scope = " AND COALESCE(pinned, 0) = 0" if exclude_pinned else "" + orphan_predicate = f"started_at < ? AND {_sql_session_last_active('sessions')} < ?" heartbeat_params: Tuple[float, ...] = () - orphan_predicate = staleness if respect_gateway_heartbeats: orphan_predicate += ( " AND NOT EXISTS (" @@ -125,14 +171,13 @@ class SessionMaintenanceMixin: " AND h.started_at <= sessions.started_at + ?" ")" ) - heartbeat_params = (hb_cutoff, hb_grace) + heartbeat_params = (now - hb_staleness, hb_grace) + scope_sql = f" AND source IN ({placeholders}){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" - f" AND source IN ({placeholders}){pin_scope}" - f" AND {orphan_predicate}", - (*srcs, cutoff, cutoff, *heartbeat_params), + f"SELECT id FROM sessions WHERE ended_at IS NULL{scope_sql}", scope_params ).fetchall() excluded = {str(x) for x in exclude_ids if x} victims = [] @@ -142,37 +187,20 @@ class SessionMaintenanceMixin: continue try: self._check_transcript_write_guards( - conn, - sid, - compression_lock_holder=None, - turn_lease_holder=None, - reject_active_turn_lease=True, - reject_active_compression_lock=True, + conn, sid, compression_lock_holder=None, turn_lease_holder=None, + reject_active_turn_lease=True, reject_active_compression_lock=True, ) - except ( - SessionCompressionInProgressError, - SessionTurnLeaseLostError, - ): + except (SessionCompressionInProgressError, SessionTurnLeaseLostError): continue victims.append(sid) if not victims: return [] - closed_at = time.time() marks = ",".join("?" for _ in victims) # 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 ({marks}) AND ended_at IS NULL" - f" AND source IN ({placeholders}){pin_scope}" - f" AND {orphan_predicate}", - ( - closed_at, - *victims, - *srcs, - cutoff, - cutoff, - *heartbeat_params, - ), + f" WHERE id IN ({marks}) AND ended_at IS NULL{scope_sql}", + (time.time(), *victims, *scope_params), ) return victims @@ -180,137 +208,30 @@ class SessionMaintenanceMixin: @staticmethod def _prune_filter_where( - *, - last_active_before: Optional[float] = None, - last_active_after: Optional[float] = None, - started_before: Optional[float] = None, - started_after: Optional[float] = None, - source: Optional[str] = None, - title_like: Optional[str] = None, - end_reason: Optional[str] = None, - cwd_prefix: Optional[str] = None, - min_messages: Optional[int] = None, - max_messages: Optional[int] = None, - archived: Optional[bool] = None, - model_like: Optional[str] = None, - provider: Optional[str] = None, - user_id: Optional[str] = None, - chat_id: Optional[str] = None, - chat_type: Optional[str] = None, - branch_like: Optional[str] = None, - min_tokens: Optional[int] = None, - max_tokens: Optional[int] = None, - min_cost: Optional[float] = None, - max_cost: Optional[float] = None, - min_tool_calls: Optional[int] = None, - max_tool_calls: Optional[int] = None, - include_pinned: bool = False, + *, archived: Optional[bool] = None, include_pinned: bool = False, **filters ) -> Tuple[str, list]: """Shared WHERE clause for bulk prune/archive selection (alias ``s``). - Filters AND together; only ended sessions are ever candidates. - ``archived`` is tri-state (None = both). ``*_like`` filters are - case-insensitive substrings; the rest are exact (provider + Filters (see ``_PRUNE_FILTERS``) AND together; only ended sessions are + ever candidates. ``archived`` is tri-state (None = both). ``*_like`` + filters are case-insensitive substrings; the rest are exact (provider case-insensitive). Token bounds use input+output; cost bounds use ``COALESCE(actual_cost_usd, estimated_cost_usd)``. """ - from hermes_state import _cwd_prefix_clause + unknown = set(filters) - _PRUNE_FILTER_NAMES + if unknown: + raise TypeError( + "SessionMaintenanceMixin._prune_filter_where() got an unexpected " + f"keyword argument {sorted(unknown)[0]!r}" + ) clauses = ["s.ended_at IS NOT NULL"] params: list = [] - if last_active_before is not None: - clauses.append( - """COALESCE( - (SELECT MAX(m.timestamp) FROM messages m - WHERE m.session_id = s.id), - s.started_at - ) < ?""" - ) - params.append(last_active_before) - # Orphan-swept rows age from the sweep, not their old activity, or - # the next prune pass deletes them before the user can recover. - clauses.append( - "(COALESCE(s.end_reason, '') != 'startup_orphan_reap' " - "OR s.ended_at < ?)" - ) - params.append(last_active_before) - if last_active_after is not None: - clauses.append( - """COALESCE( - (SELECT MAX(m.timestamp) FROM messages m - WHERE m.session_id = s.id), - s.started_at - ) >= ?""" - ) - params.append(last_active_after) - if started_before is not None: - clauses.append("s.started_at < ?") - params.append(started_before) - if started_after is not None: - clauses.append("s.started_at >= ?") - params.append(started_after) - if source: - clauses.append("s.source = ?") - params.append(source) - if title_like: - clauses.append("LOWER(COALESCE(s.title, '')) LIKE ? ESCAPE '\\'") - params.append(f"%{_escape_like(title_like.lower())}%") - if end_reason: - clauses.append("s.end_reason = ?") - params.append(end_reason) - if cwd_prefix: - clause, clause_params = _cwd_prefix_clause(cwd_prefix) - clauses.append(clause) - params.extend(clause_params) - if min_messages is not None: - clauses.append("s.message_count >= ?") - params.append(min_messages) - if max_messages is not None: - clauses.append("s.message_count <= ?") - params.append(max_messages) - if model_like: - clauses.append("LOWER(COALESCE(s.model, '')) LIKE ? ESCAPE '\\'") - params.append(f"%{_escape_like(model_like.lower())}%") - if provider: - clauses.append("LOWER(COALESCE(s.billing_provider, '')) = ?") - params.append(provider.lower()) - if user_id: - clauses.append("s.user_id = ?") - params.append(user_id) - if chat_id: - clauses.append("s.chat_id = ?") - params.append(chat_id) - if chat_type: - clauses.append("s.chat_type = ?") - params.append(chat_type) - if branch_like: - clauses.append("LOWER(COALESCE(s.git_branch, '')) LIKE ? ESCAPE '\\'") - params.append(f"%{_escape_like(branch_like.lower())}%") - if min_tokens is not None: - clauses.append( - "(COALESCE(s.input_tokens, 0) + COALESCE(s.output_tokens, 0)) >= ?" - ) - params.append(min_tokens) - if max_tokens is not None: - clauses.append( - "(COALESCE(s.input_tokens, 0) + COALESCE(s.output_tokens, 0)) <= ?" - ) - params.append(max_tokens) - if min_cost is not None: - clauses.append( - "COALESCE(s.actual_cost_usd, s.estimated_cost_usd, 0) >= ?" - ) - params.append(min_cost) - if max_cost is not None: - clauses.append( - "COALESCE(s.actual_cost_usd, s.estimated_cost_usd, 0) <= ?" - ) - params.append(max_cost) - if min_tool_calls is not None: - clauses.append("COALESCE(s.tool_call_count, 0) >= ?") - params.append(min_tool_calls) - if max_tool_calls is not None: - clauses.append("COALESCE(s.tool_call_count, 0) <= ?") - params.append(max_tool_calls) + for name, applies, build in _PRUNE_FILTERS: + value = filters.get(name) + if (value is not None) if applies == "notnone" else bool(value): + new_clauses, new_params = build(value) + clauses.extend(new_clauses) + params.extend(new_params) if archived is True: clauses.append("s.archived = 1") elif archived is False: @@ -322,33 +243,28 @@ class SessionMaintenanceMixin: return " AND ".join(clauses), params @staticmethod - def _apply_prune_age_filter( - older_than_days: Optional[float], filters: Dict[str, Any] - ) -> None: + def _apply_prune_age_filter(older_than_days: Optional[float], filters: Dict[str, Any]) -> None: """Translate the legacy age window into the shared activity filter.""" if ( filters.get("last_active_before") is None and filters.get("started_before") is None and older_than_days is not None ): - filters["last_active_before"] = time.time() - ( - older_than_days * 86400 - ) + filters["last_active_before"] = time.time() - (older_than_days * 86400) + + def _prune_where(self, older_than_days, source, filters) -> Tuple[str, list]: + self._apply_prune_age_filter(older_than_days, filters) + return self._prune_filter_where(source=source, **filters) def list_prune_candidates( - self, - older_than_days: Optional[float] = None, - source: str = None, - **filters, + self, older_than_days: Optional[float] = None, source: str = None, **filters ) -> List[Dict[str, Any]]: """Sessions a matching prune/archive would touch (dry-run), oldest first. Same filters as :meth:`_prune_filter_where`; ``older_than_days`` is an inactivity threshold (latest message, else ``started_at``).""" - self._apply_prune_age_filter(older_than_days, filters) - where, params = self._prune_filter_where(source=source, **filters) - with self._read_ctx() as conn: - cursor = conn.execute( - f"""SELECT s.id, s.source, s.title, s.model, s.started_at, + where, params = self._prune_where(older_than_days, source, filters) + rows = self._read_all( + f"""SELECT s.id, s.source, s.title, s.model, s.started_at, COALESCE( (SELECT MAX(m.timestamp) FROM messages m WHERE m.session_id = s.id), @@ -357,50 +273,32 @@ class SessionMaintenanceMixin: s.ended_at, s.message_count, s.archived FROM sessions s WHERE {where} ORDER BY last_active ASC, s.started_at ASC""", - params, - ) - return [dict(row) for row in cursor.fetchall()] + params, + ) + return [dict(row) for row in rows] def count_prune_matches( - self, - older_than_days: Optional[float] = None, - source: str = None, - **filters, + self, older_than_days: Optional[float] = None, source: str = None, **filters ) -> int: """Count-only variant of :meth:`list_prune_candidates` (the CLI uses it to report how many pinned sessions are spared).""" - self._apply_prune_age_filter(older_than_days, filters) - where, params = self._prune_filter_where(source=source, **filters) - with self._read_ctx() as conn: - cursor = conn.execute( - f"SELECT COUNT(*) FROM sessions s WHERE {where}", params - ) - return int(cursor.fetchone()[0]) + where, params = self._prune_where(older_than_days, source, filters) + return int(self._read_one(f"SELECT COUNT(*) FROM sessions s WHERE {where}", params)[0]) def count_open_prune_matches( - self, - older_than_days: Optional[float] = None, - source: str = None, - **filters, + self, older_than_days: Optional[float] = None, source: str = None, **filters ) -> int: """Count open sessions a matching prune skips: every normal filter with only the ``ended_at`` guard inverted. Visibility-only; live sessions never become prune-eligible.""" - self._apply_prune_age_filter(older_than_days, filters) - where, params = self._prune_filter_where(source=source, **filters) + where, params = self._prune_where(older_than_days, source, filters) ended_guard = "s.ended_at IS NOT NULL" if not where.startswith(ended_guard): raise RuntimeError("prune filter lost its ended-session safety guard") open_where = f"s.ended_at IS NULL{where[len(ended_guard):]}" - with self._read_ctx() as conn: - cursor = conn.execute( - f"SELECT COUNT(*) FROM sessions s WHERE {open_where}", params - ) - return int(cursor.fetchone()[0]) + return int(self._read_one(f"SELECT COUNT(*) FROM sessions s WHERE {open_where}", params)[0]) - def archive_stale_sessions( - self, idle_days: float, *, exclude_pinned: bool = True - ) -> int: + def archive_stale_sessions(self, idle_days: float, *, exclude_pinned: bool = True) -> int: """Archive every session untouched for ``idle_days`` (real recency: freshest of ``last_activity_at`` / latest message / ``started_at``). Unlike :meth:`archive_sessions`, this can archive unended sessions. @@ -432,11 +330,8 @@ class SessionMaintenanceMixin: return len(ids) def prune_sessions( - self, - older_than_days: Optional[float] = 90, - source: str = None, - sessions_dir: Optional[Path] = None, - exclude_active_write_guards: bool = False, + self, older_than_days: Optional[float] = 90, source: str = None, + sessions_dir: Optional[Path] = None, exclude_active_write_guards: bool = False, **filters, ) -> int: """Delete ended sessions matching the filters; returns the count. @@ -454,46 +349,32 @@ class SessionMaintenanceMixin: while expired/dead holders are reclaimed and fenced in the same write. """ from hermes_state import SessionCompressionInProgressError, SessionTurnLeaseLostError - self._apply_prune_age_filter(older_than_days, filters) - where, where_params = self._prune_filter_where(source=source, **filters) + where, where_params = self._prune_where(older_than_days, source, filters) removed_ids: list[str] = [] def _do(conn): - cursor = conn.execute( - f"SELECT s.id FROM sessions s WHERE {where}", where_params - ) + cursor = conn.execute(f"SELECT s.id FROM sessions s WHERE {where}", where_params) session_ids = {row["id"] for row in cursor.fetchall()} - if exclude_active_write_guards: protected = set() for sid in session_ids: try: self._check_transcript_write_guards( - conn, - sid, - compression_lock_holder=None, - turn_lease_holder=None, - reject_active_turn_lease=True, - reject_active_compression_lock=True, + conn, sid, compression_lock_holder=None, turn_lease_holder=None, + reject_active_turn_lease=True, reject_active_compression_lock=True, allow_closed_compression_parent=True, ) - except ( - SessionCompressionInProgressError, - SessionTurnLeaseLostError, - ): + except (SessionCompressionInProgressError, SessionTurnLeaseLostError): protected.add(sid) session_ids.difference_update(protected) - if not session_ids: return 0 - placeholders = ",".join("?" * len(session_ids)) conn.execute( f"UPDATE sessions SET parent_session_id = NULL " f"WHERE parent_session_id IN ({placeholders})", 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,)) @@ -506,6 +387,14 @@ class SessionMaintenanceMixin: self._remove_session_files(sessions_dir, sid) return count + def _page_pragmas(self, *names: str) -> Optional[list]: + """Read integer PRAGMAs over the existing connection (never a byte probe + of the live file); None if the connection is closed or a pragma fails.""" + with self._read_ctx() as conn: + if self._conn is None: + return None + return [conn.execute(f"PRAGMA {name}").fetchone()[0] for name in names] + def logical_size_bytes(self) -> Optional[int]: """``page_count * page_size``: the main-file size once the WAL is checkpointed back in. Prefer over ``os.path.getsize`` when reporting a @@ -514,28 +403,24 @@ class SessionMaintenanceMixin: understates the win and can go negative. None if pragmas fail. """ try: - with self._read_ctx() as conn: - if self._conn is None: - return None - page_count = conn.execute("PRAGMA page_count").fetchone()[0] - page_size = conn.execute("PRAGMA page_size").fetchone()[0] + values = self._page_pragmas("page_count", "page_size") + if values is None: + return None + page_count, page_size = values return int(page_count) * int(page_size) except Exception as exc: logger.debug("Could not read logical DB size: %s", exc) return None def _freelist_ratio(self) -> Optional[float]: - """Reclaimable fraction (``freelist_count / page_count``) over the - existing connection — never a byte-level probe of the live file. Gates - VACUUM in :meth:`maybe_auto_prune_and_vacuum`. None if pragmas fail - (callers then fall back to the time throttle alone). - """ + """Reclaimable fraction (``freelist_count / page_count``); gates VACUUM + in :meth:`maybe_auto_prune_and_vacuum`. None if pragmas fail (callers + then fall back to the time throttle alone).""" try: - with self._read_ctx() as conn: - if self._conn is None: - return None - page_count = int(conn.execute("PRAGMA page_count").fetchone()[0]) - freelist = int(conn.execute("PRAGMA freelist_count").fetchone()[0]) + values = self._page_pragmas("page_count", "freelist_count") + if values is None: + return None + page_count, freelist = int(values[0]), int(values[1]) if page_count <= 0: return 0.0 return freelist / page_count @@ -553,13 +438,11 @@ class SessionMaintenanceMixin: :meth:`optimize_fts` so the VACUUM reclaims those pages too. Returns the number of FTS indexes optimized (0 on merge failure / no FTS). """ - # optimize_fts() manages its own lock. optimized = 0 try: - optimized = self.optimize_fts() + optimized = self.optimize_fts() # manages its own lock except Exception as exc: logger.warning("FTS optimize before VACUUM failed: %s", exc) - # VACUUM cannot be executed inside a transaction. with self._lock: # PASSIVE, not TRUNCATE: a manual `hermes sessions vacuum` runs in # a transient CLI process, and a TRUNCATE reset here would race a @@ -570,8 +453,7 @@ class SessionMaintenanceMixin: logger.debug("WAL checkpoint (PASSIVE) before VACUUM failed: %s", exc) self._conn.execute("VACUUM") # VACUUM rewrites every page THROUGH the WAL; without this TRUNCATE - # a 3 GB database leaves a 3 GB -wal behind and the command is a - # net loss on disk. + # a 3 GB database leaves a 3 GB -wal behind. try: self._conn.execute("PRAGMA wal_checkpoint(TRUNCATE)") except Exception as exc: @@ -582,12 +464,8 @@ class SessionMaintenanceMixin: return optimized def maybe_auto_prune_and_vacuum( - self, - retention_days: int = 90, - min_interval_hours: int = 24, - vacuum: bool = True, - sessions_dir: Optional[Path] = None, - min_vacuum_interval_days: int = 30, + self, retention_days: int = 90, min_interval_hours: int = 24, vacuum: bool = True, + sessions_dir: Optional[Path] = None, min_vacuum_interval_days: int = 30, min_vacuum_freelist_ratio: float = AUTO_VACUUM_MIN_FREELIST_RATIO, ) -> Dict[str, Any]: """Idempotent startup auto-maintenance: prune inactive sessions, @@ -611,12 +489,7 @@ class SessionMaintenanceMixin: failure. """ from hermes_state import _release_auto_maintenance_lock, _try_acquire_auto_maintenance_lock - result: Dict[str, Any] = { - "skipped": False, - "pruned": 0, - "closed": 0, - "vacuumed": False, - } + result: Dict[str, Any] = {"skipped": False, "pruned": 0, "closed": 0, "vacuumed": False} maintenance_lock = _try_acquire_auto_maintenance_lock(self.db_path) if maintenance_lock is None: result["skipped"] = True @@ -626,8 +499,7 @@ class SessionMaintenanceMixin: now = time.time() if last_raw: try: - last_ts = float(last_raw) - if now - last_ts < min_interval_hours * 3600: + if now - float(last_raw) < min_interval_hours * 3600: result["skipped"] = True return result except (TypeError, ValueError): @@ -635,23 +507,18 @@ class SessionMaintenanceMixin: # Prune first: orphans closed below get a full retention window. pruned = self.prune_sessions( - older_than_days=retention_days, - sessions_dir=sessions_dir, + older_than_days=retention_days, sessions_dir=sessions_dir, exclude_active_write_guards=True, ) result["pruned"] = pruned - closed = self.sweep_orphaned_sessions( max_idle_seconds=float(retention_days) * 86400.0, - sources=self._AUTO_PRUNE_STALE_OPEN_SOURCES, - exclude_pinned=True, - # State-owned lifecycles, not gateway heartbeats. - respect_gateway_heartbeats=False, + sources=self._AUTO_PRUNE_STALE_OPEN_SOURCES, exclude_pinned=True, + respect_gateway_heartbeats=False, # state-owned lifecycles, not gateway heartbeats ) result["closed"] = len(closed) - # VACUUM only if rows were freed, the time throttle passed ("not - # too often") AND the freelist ratio passed ("only when it pays - # off") — it holds an exclusive lock for a full rewrite. + # VACUUM only if rows were freed, the time throttle passed AND the + # freelist ratio passed — it holds an exclusive lock for a full rewrite. last_vacuum_raw = self.get_meta("last_vacuum") vacuum_due = True if last_vacuum_raw: @@ -673,21 +540,15 @@ class SessionMaintenanceMixin: logger.debug( "state.db auto-maintenance: skipping VACUUM, only " "%.1f%% of pages reclaimable (threshold %.0f%%)", - ratio * 100.0, - min_vacuum_freelist_ratio * 100.0, + ratio * 100.0, min_vacuum_freelist_ratio * 100.0, ) - # Record even when pruned == 0 so the throttle holds. self.set_meta("last_auto_prune", str(now)) - if closed or pruned > 0: logger.info( "state.db auto-maintenance: closed %d stale open session(s), " "pruned %d session(s) inactive for %d days%s", - len(closed), - pruned, - retention_days, - " + VACUUM" if result["vacuumed"] else "", + len(closed), pruned, retention_days, " + VACUUM" if result["vacuumed"] else "", ) except Exception as exc: # Maintenance must never block startup. @@ -695,5 +556,4 @@ class SessionMaintenanceMixin: result["error"] = str(exc) finally: _release_auto_maintenance_lock(maintenance_lock) - return result