From af85e7c6446f72d4c5c456dccf5bf28da4448090 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 19:21:56 -0700 Subject: [PATCH] =?UTF-8?q?refactor(state):=20compact=20SessionGatewayMixi?= =?UTF-8?q?n/SessionCompressionMixin=20=E2=80=94=20compose=20fail=5Fhandof?= =?UTF-8?q?f/lineage=20SQL,=20unify=20set=5F*=20writers,=20trim=20docstrin?= =?UTF-8?q?gs?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- hermes_state_compression.py | 212 ++++++++--------- hermes_state_gateway.py | 462 ++++++++++++------------------------ 2 files changed, 251 insertions(+), 423 deletions(-) diff --git a/hermes_state_compression.py b/hermes_state_compression.py index 1110b66880..6f06ee35a0 100644 --- a/hermes_state_compression.py +++ b/hermes_state_compression.py @@ -28,10 +28,8 @@ def _ended_by_compression(row) -> bool: 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} + return {"session_exists": exists, + "cooldown_until": float(cooldown_until) if cooldown_until is not None else None, "error": error} def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float, @@ -55,10 +53,9 @@ class SessionCompressionMixin: """Compression lineage, cooldown/streak counters, locks and turn leases.""" def find_live_compression_child(self, parent_session_id: str) -> Optional[Dict[str, Any]]: - """The unique live direct child of a compression-ended session, else None. - - A stale agent whose parent was rotated elsewhere may recover only when the - lineage names exactly one live continuation; more than one fails closed.""" + """The unique live direct child of a compression-ended session, else None. A stale + agent whose parent was rotated elsewhere may recover only when the lineage names + exactly one live continuation; more than one fails closed.""" if not parent_session_id: return None with self._read_ctx() as conn: @@ -108,9 +105,9 @@ class SessionCompressionMixin: ).fetchone() if child is not None: return False - # refresh_compression_lock() lets an owner revive its own expired row, so - # reclaim it inside this write txn: refresh-first makes the lease active and - # aborts recovery; recovery-first deletes the holder so a refresh can't resurrect it. + # refresh_compression_lock() lets an owner revive its own expired row, so reclaim + # it inside this write txn: refresh-first makes the lease active and aborts + # recovery; recovery-first deletes the holder so a refresh can't resurrect it. now = time.time() lock_row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone() if lock_row is not None: @@ -118,8 +115,7 @@ class SessionCompressionMixin: if expires_at is None or float(expires_at) >= now: return False deleted = conn.execute( - "DELETE FROM compression_locks " - "WHERE session_id = ? AND holder = ? AND expires_at = ?", + "DELETE FROM compression_locks WHERE session_id = ? AND holder = ? AND expires_at = ?", (session_id, lock_row["holder"], expires_at)) if deleted.rowcount != 1: return False @@ -127,9 +123,9 @@ class SessionCompressionMixin: "UPDATE sessions SET ended_at = NULL, end_reason = NULL " "WHERE id = ? AND ended_at IS NOT NULL AND end_reason = 'compression'", (session_id,)) - # rowcount==1 is guaranteed by the parent SELECT in this same txn. A False - # return added past this point must raise instead: the lease DELETE above - # commits unless _do raises. + # rowcount==1 is guaranteed by the parent SELECT in this same txn. A False return + # 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)) @@ -164,41 +160,35 @@ class SessionCompressionMixin: system_prompt: str = None, cwd: str = None, profile_name: str = None, compression_lock_holder: str = None, require_compression_lease: bool = True, require_lease_refresh: bool = False, lease_ttl_seconds: float = 300.0, - watermark: Optional[int] = None, watermark_ceiling: Optional[int] = None, - ) -> None: - """Atomically close a parent and publish its durable compression child. - - Closure, child row, and handoff commit in one transaction: readers see the live - parent or a complete child, never an ended parent with a missing/empty child. + watermark: Optional[int] = None, watermark_ceiling: Optional[int] = None) -> None: + """Atomically close a parent and publish its durable compression child: closure, + child row, and handoff commit in one transaction, so readers see the live parent + or a complete child, never an ended parent with a missing/empty child. *watermark* (parent's ``get_active_message_watermark`` at compression start): - parent rows with ``id > watermark`` — appends landed during the slow summary — - are column-cloned into the child AFTER the handoff. *watermark_ceiling* bounds - the clone: the rotation path flushes its OWN transcript to the parent just - before publishing and those rows are already in the handoff, so only + parent rows with ``id > watermark`` — appends landed during the slow summary — are + column-cloned into the child AFTER the handoff. *watermark_ceiling* bounds the + clone: the rotation path flushes its OWN transcript to the parent just before + publishing and those rows are already in the handoff, so only ``(watermark, watermark_ceiling]`` is foreign tail (``None`` = unbounded). - - *require_lease_refresh* + *compression_lock_holder* refreshes the lease on the - same ``conn`` before the expiry check (no TOCTOU window), so a refresher that - died on transient DB errors gets one last chance.""" + *require_lease_refresh* + *compression_lock_holder* refreshes the lease on the same + ``conn`` before the expiry check (no 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( - "UPDATE compression_locks SET expires_at = ? " - "WHERE session_id = ? AND holder = ?", + "UPDATE compression_locks SET expires_at = ? WHERE session_id = ? AND holder = ?", (time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder)) lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone() if require_compression_lease and ( - lock_row is None - or not compression_lock_holder + lock_row is None or not compression_lock_holder or lock_row["holder"] != compression_lock_holder or float(lock_row["expires_at"]) <= time.time() ): raise CompressionSessionBusyError( - f"Compression lease lost before publication: {parent_session_id}" - ) + f"Compression lease lost before publication: {parent_session_id}") parent = conn.execute( """SELECT ended_at, end_reason, cwd, git_branch, git_repo_root, user_id, session_key, chat_id, chat_type, @@ -209,17 +199,15 @@ class SessionCompressionMixin: if parent is None: raise RuntimeError(f"Compression parent not found: {parent_session_id}") if parent["ended_at"] is not None: - # An AUTOMATIC end stamp (tui_shutdown, ws_disconnect, orphan reap, - # idle/LRU evict) is stale by construction — this lease holder is still - # continuing the conversation, and left alone it wedges rotation forever. - # Clear it; the closure UPDATE below re-stamps end_reason='compression'. - # Deliberate boundaries still fail closed. - if is_automatic_end_reason(parent["end_reason"]): - conn.execute( - "UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?", - (parent_session_id,)) - else: + # An AUTOMATIC end stamp (tui_shutdown, ws_disconnect, orphan reap, idle/LRU + # evict) is stale by construction — this lease holder is still continuing the + # conversation, and left alone it wedges rotation forever. Clear it; the closure + # UPDATE below re-stamps end_reason='compression'. Deliberate boundaries fail closed. + if not is_automatic_end_reason(parent["end_reason"]): raise RuntimeError(f"Compression parent already ended: {parent_session_id}") + conn.execute( + "UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?", + (parent_session_id,)) if not messages: raise RuntimeError("Compression child handoff must not be empty") self._publish_child_session_row( @@ -235,8 +223,7 @@ class SessionCompressionMixin: conn, "SELECT id, tool_calls FROM messages " "WHERE session_id = ? AND active = 1 AND id > ?" f"{' AND id <= ?' if bounded else ''} ORDER BY id", - [parent_session_id, int(watermark), *([int(watermark_ceiling)] if bounded else [])], - ) + [parent_session_id, int(watermark), *([int(watermark_ceiling)] if bounded else [])]) if tail_ids: self._clone_message_rows(conn, tail_ids, session_id=child_session_id) total_messages += len(tail_ids) @@ -259,19 +246,18 @@ class SessionCompressionMixin: except sqlite3.Error as exc: logger.warning("%s(%s) failed: %s", op, session_id, exc) - def record_compression_failure_cooldown(self, session_id: str, cooldown_until: float, error: Optional[str] = None) -> None: - """Persist the active compression-failure cooldown. Merge-max with any longer - live deadline so a later shorter write can't reopen the thrash window; error - always takes the latest diagnostic.""" + def record_compression_failure_cooldown( + self, session_id: str, cooldown_until: float, error: Optional[str] = None) -> None: + """Persist the active compression-failure cooldown. Merge-max with any longer live + deadline so a later shorter write can't reopen the thrash window; error always + takes the latest diagnostic.""" if not session_id: return self._write_sql_logged( "record_compression_failure_cooldown", session_id, "UPDATE sessions SET compression_failure_cooldown_until = CASE " - "WHEN compression_failure_cooldown_until IS NOT NULL " - " AND compression_failure_cooldown_until > ? " - "THEN compression_failure_cooldown_until ELSE ? END, " - "compression_failure_error = ? WHERE id = ?", + "WHEN compression_failure_cooldown_until IS NOT NULL AND compression_failure_cooldown_until > ? " + "THEN compression_failure_cooldown_until ELSE ? END, compression_failure_error = ? WHERE id = ?", (cooldown_until, cooldown_until, error, session_id)) def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]: @@ -285,17 +271,15 @@ class SessionCompressionMixin: return {"cooldown_until": float(row[0]), "remaining_seconds": float(row[0]) - now, "error": row[1]} def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]: - """Exact stored cooldown columns, no expiry filtering, so compression - cancellation can roll back an expired, partially-null, or absent row exactly.""" + """Exact stored cooldown columns, no expiry filtering, so compression cancellation + can roll back an expired, partially-null, or absent row exactly.""" row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None - if row is None: - return _cooldown_row(False, None, None) - return _cooldown_row(True, row[0], row[1]) + return _cooldown_row(False, None, None) if row is None else _cooldown_row(True, row[0], row[1]) def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None: """Restore and verify an exact cooldown-row snapshot. Unlike record/clear this - rollback API propagates write and verification failures: cancellation must not - be reported mutation-free when compensation failed.""" + rollback API propagates write and verification failures: cancellation must not be + reported mutation-free when compensation failed.""" if not snapshot.get("session_exists", False): if self.get_compression_failure_cooldown_row(session_id).get("session_exists", False): raise RuntimeError("cannot restore absent compression cooldown row: session now exists") @@ -340,6 +324,9 @@ class SessionCompressionMixin: except (TypeError, ValueError): return zero + def _write_session_column(self, column: str, session_id: str, value: Any) -> None: + self._write_sql(f"UPDATE sessions SET {column} = ? WHERE id = ?", (value, session_id)) + def get_compression_fallback_streak(self, session_id: str) -> int: """Return the persisted deterministic-fallback streak.""" return self._read_session_number("compression_fallback_streak", session_id, int, 0) @@ -347,22 +334,18 @@ class SessionCompressionMixin: def set_compression_fallback_streak(self, session_id: str, streak: int) -> None: """Persist the deterministic-fallback streak for one session.""" if session_id: - self._write_sql( - "UPDATE sessions SET compression_fallback_streak = ? WHERE id = ?", - (max(0, int(streak)), session_id)) + self._write_session_column("compression_fallback_streak", session_id, max(0, int(streak))) def get_compression_ineffective_count(self, session_id: str) -> int: - """Persisted ineffective-compaction strike count — the durable half of the - built-in compressor's anti-thrash guard, so a fresh compressor bound to a resumed - session inherits an armed/tripped guard across restarts.""" + """Persisted ineffective-compaction strike count — the durable half of the built-in + compressor's anti-thrash guard, so a fresh compressor bound to a resumed session + inherits an armed/tripped guard across restarts.""" return self._read_session_number("compression_ineffective_count", session_id, int, 0) def set_compression_ineffective_count(self, session_id: str, count: int) -> None: """Persist the ineffective-compaction strike count for one session.""" if session_id: - self._write_sql( - "UPDATE sessions SET compression_ineffective_count = ? WHERE id = ?", - (max(0, int(count)), session_id)) + self._write_session_column("compression_ineffective_count", session_id, max(0, int(count))) def get_compression_recovery_deadline(self, session_id: str) -> float: """Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable @@ -377,19 +360,17 @@ class SessionCompressionMixin: normalized = max(0.0, float(deadline or 0.0)) except (TypeError, ValueError): normalized = 0.0 - self._write_sql( - "UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?", - (normalized or None, session_id)) + self._write_session_column("compression_recovery_deadline", session_id, normalized or None) def refresh_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool: """Extend the compression lock lease if ``holder`` still owns it. Ownership is decided by ``holder`` alone, deliberately NOT ``expires_at``: a live owner whose refresher stalled past its TTL must be able to revive its still- - unclaimed row, otherwise it keeps compressing with no lease — the window in - which a competing path can fork the lineage. It cannot resurrect a lock someone - else took: SQLite serialises writes, so the reclaim (DELETE-expired + INSERT OR - IGNORE) never interleaves with this UPDATE.""" + unclaimed row, otherwise it keeps compressing with no lease — the window in which + a competing path can fork the lineage. It cannot resurrect a lock someone else + took: SQLite serialises writes, so the reclaim (DELETE-expired + INSERT OR IGNORE) + never interleaves with this UPDATE.""" if not session_id or not holder: return False expires_at = time.time() + ttl_seconds @@ -403,13 +384,10 @@ class SessionCompressionMixin: return False def try_acquire_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool: - """Try to atomically acquire the compression lock for ``session_id``. - - ``False``: another holder owns a live lock and the caller MUST NOT compress (its - rotation would split the lineage). Expired locks and structured holders whose - local ``pid=`` is dead are reclaimed transparently. Single-transaction DELETE- - expired + INSERT OR IGNORE + SELECT-to-confirm (INSERT OR IGNORE gives no - rowcount signal).""" + """Try to atomically acquire the compression lock for ``session_id``. ``False``: + another holder owns a live lock and the caller MUST NOT compress (its rotation + would split the lineage). Expired locks and structured holders whose local ``pid=`` + is dead are reclaimed transparently.""" from hermes_state import _compression_lock_holder_process_is_dead if not session_id: return False @@ -424,7 +402,8 @@ class SessionCompressionMixin: try: acquired, reclaimed_holder = self._execute_write(_do) if reclaimed_holder: - logger.warning("Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder) + logger.warning("Reclaimed stale compression lock for session=%s (holder=%s)", + session_id, reclaimed_holder) return bool(acquired) except sqlite3.Error as exc: # False makes the caller skip compression — safe when the lock subsystem is broken. @@ -441,18 +420,17 @@ class SessionCompressionMixin: (session_id, holder)) def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str: - """Walk compression parents on ``conn`` to the conversation lease key. - - Must share the connection of the lease INSERT/UPDATE/DELETE: a failed lookup - must not yield a child id the write then persists. Markers bind to - ``parent_session_id``. Lock errors propagate so ``_execute_write`` can retry.""" + """Walk compression parents on ``conn`` to the conversation lease key. Must share + the connection of the lease INSERT/UPDATE/DELETE: a failed lookup must not yield a + child id the write then persists. Markers bind to ``parent_session_id``. Lock + 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 = ?", (sid,), - ).fetchone() + "SELECT id, parent_session_id, source, model_config, end_reason FROM sessions WHERE id = ?", + (sid,)).fetchone() return dict(row) if row else None current = _row(session_id) @@ -469,8 +447,8 @@ class SessionCompressionMixin: return str(current.get("id") or session_id) if current else session_id def _session_turn_lease_key(self, session_id: str) -> str: - """Stable serialization key for every compression segment (tests/diagnostics; - the write paths resolve it inside their own txn). Does not swallow lock errors.""" + """Stable serialization key for every compression segment (tests/diagnostics; the + write paths resolve it inside their own txn). Does not swallow lock errors.""" if not session_id: return session_id with self._read_ctx() as conn: @@ -479,9 +457,9 @@ class SessionCompressionMixin: def try_acquire_session_turn_lease( self, session_id: str, holder: str, *, ttl_seconds: float = 300.0, patience_s: Optional[float] = None, ) -> bool: - """Atomically acquire the cross-process turn lease for a conversation (keyed by - the lineage root). The walk, the INSERT, and reclaim of expired or dead-local-PID - leases share one write transaction.""" + """Atomically acquire the cross-process turn lease for a conversation (keyed by the + lineage root). The walk, the INSERT, and reclaim of expired or dead-local-PID leases + share one write transaction.""" from hermes_state import _compression_lock_holder_process_is_dead if not session_id or not holder: return False @@ -500,14 +478,12 @@ class SessionCompressionMixin: def acquire_session_turn_lease( self, session_id: str, holder: str, *, ttl_seconds: float = 300.0, wait_seconds: float = 1800.0, poll_interval_seconds: float = 1.0, on_wait=None, - wait_notice_interval_seconds: float = 15.0, should_abort=None, - acquire_patience_s: float = 0.5, + wait_notice_interval_seconds: float = 15.0, should_abort=None, acquire_patience_s: float = 0.5, ) -> bool: - """Wait for a cross-process turn lease without holding a SQLite lock. - - ``on_wait(elapsed)`` is best-effort: called when the first attempt fails and - about every ``wait_notice_interval_seconds`` after. ``should_abort()`` True - (e.g. ``/stop``) returns False at once.""" + """Wait for a cross-process turn lease without holding a SQLite lock. ``on_wait(elapsed)`` + is best-effort: called when the first attempt fails and about every + ``wait_notice_interval_seconds`` after. ``should_abort()`` True (e.g. ``/stop``) returns + False at once.""" from hermes_state import classify_persistence_error deadline = time.monotonic() + max(0.0, float(wait_seconds)) wait_started = None @@ -583,8 +559,8 @@ class SessionCompressionMixin: return None if row is None else row[0] def finalize_orphaned_compression_sessions(self) -> int: - """Mark orphaned compression continuations (parent ended by compression; child - has messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as + """Mark orphaned compression continuations (parent ended by compression; child has + messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as ``orphaned_compression``. Non-destructive.""" cutoff = time.time() - 604800 # 7 days return self._write_rowcount( @@ -613,15 +589,13 @@ class SessionCompressionMixin: def get_compression_chain(self, session_id: str) -> List[str]: """Walk the compression-continuation chain forward: root-first through the tip - (``[session_id]`` when no continuation). ``get_compression_tip`` is this walk's - last element. - - A continuation is a child of a session with ``end_reason='compression'``. The - old ``child.started_at >= parent.ended_at`` test was too brittle (gateway + - compression races insert the real continuation before ``ended_at`` is written, - while a stale websocket later creates a sibling that passes it). Instead exclude - branch/delegate/tool children and prefer children that continue the chain or - are still live over stale closed siblings such as ``ws_orphan_reap``.""" + (``[session_id]`` when no continuation); ``get_compression_tip`` is the last element. + A continuation is a child of a session with ``end_reason='compression'``. The old + ``child.started_at >= parent.ended_at`` test was too brittle (gateway + compression + races insert the real continuation before ``ended_at`` is written, while a stale + websocket later creates a sibling that passes it). Instead exclude + branch/delegate/tool children and prefer children that continue the chain or are + still live over stale closed siblings such as ``ws_orphan_reap``.""" current = session_id chain = [current] if current else [] seen = set(chain) @@ -659,8 +633,8 @@ class SessionCompressionMixin: return chain def get_compression_tip(self, session_id: str) -> Optional[str]: - """Live tip of a compression chain (``get_compression_chain`` semantics); the - input id when no continuation exists.""" + """Live tip of a compression chain (``get_compression_chain`` semantics); the input + id when no continuation exists.""" chain = self.get_compression_chain(session_id) return chain[-1] if chain else session_id diff --git a/hermes_state_gateway.py b/hermes_state_gateway.py index a54e305052..883e4d6130 100644 --- a/hermes_state_gateway.py +++ b/hermes_state_gateway.py @@ -1,8 +1,7 @@ """Gateway-facing SessionDB persistence: routing index, peers, orphans, heartbeats, handoffs. Mixin bound onto ``SessionDB`` via the MRO; built on its ``_read_ctx`` / -``_execute_write`` / ``_write_sql`` / ``_read_all`` primitives. -""" +``_execute_write`` / ``_write_sql`` / ``_read_all`` primitives.""" from __future__ import annotations @@ -13,11 +12,7 @@ import time from pathlib import Path from typing import Any, Dict, List, Optional, Set, Tuple -from hermes_state_common import ( - _RECOVERABLE_END_REASONS_SQL, - _RESET_END_REASONS_SQL, - _sql_session_last_active, -) +from hermes_state_common import _RECOVERABLE_END_REASONS_SQL, _RESET_END_REASONS_SQL, _sql_session_last_active # Log-record parity with the origin module (caplog tests pin "hermes_state"). logger = logging.getLogger("hermes_state") @@ -140,19 +135,18 @@ _ORPHAN_CONTIGUITY_DONORS_SQL = f""" ORDER BY last_active DESC LIMIT 2 """ +_HANDOFF_FAIL_SQL = "UPDATE sessions SET handoff_state = 'failed', handoff_error = ? WHERE " class SessionGatewayMixin: """Routing index, session peers/orphans, hygiene streaks, heartbeats, handoffs.""" def _reap_inactive_orphan_desktop_holders( - self, holders: List[Tuple[int, str]], *, min_age_seconds: float - ) -> List[int]: + self, holders: List[Tuple[int, str]], *, min_age_seconds: float) -> List[int]: """Terminate old PPID-1 Desktop ephemeral backends with no client. Fails closed: anything whose parent, age, argv, or network connections - cannot be proved safe remains a repair-blocking holder. - """ + cannot be proved safe remains a repair-blocking holder.""" from hermes_state import _concrete_state_db_holder_pids, _is_inactive_orphan_desktop_holder, psutil if not sys.platform.startswith("linux") or psutil is None: return [] @@ -160,7 +154,6 @@ class SessionGatewayMixin: from hermes_cli.dashboard_procs import _is_ephemeral_port_zero_backend except Exception: return [] - now = time.time() candidates = [] for pid in _concrete_state_db_holder_pids(self.db_path, holders): @@ -171,13 +164,11 @@ class SessionGatewayMixin: 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, - ): + connection_statuses=statuses): continue except Exception: continue candidates.append(process) - signalled: List[int] = [] for process in candidates: try: @@ -205,54 +196,40 @@ class SessionGatewayMixin: 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, - include_compression_ancestors: bool = False, - ) -> 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. - ``display_name`` / ``origin_json`` let consumers (mcp_serve, mirror, - channel directory) read routing data from state.db instead of - sessions.json; ``None`` leaves the existing value untouched. - ``include_compression_ancestors`` keeps a compression lineage on one - routing peer when an explicit resume moves its tip to another lane; - normal per-turn refreshes update only the supplied row. - - Self-healing: a missing target row (deferred ``create_session`` write, - or crash between routing publication and row creation) is INSERTed - with full identity rather than silently no-opped, so a gateway row can - never be first-created by an identity-less lazy writer - (``update_token_counts``) and stay unroutable forever. - """ + ``display_name`` / ``origin_json``: ``None`` leaves the stored value untouched + (consumers read routing data from state.db, not sessions.json). + ``include_compression_ancestors`` keeps a compression lineage on one routing + peer when an explicit resume moves its tip to another lane; per-turn refreshes + update only the supplied row. Self-healing: a missing target row (deferred + ``create_session`` write, or crash between routing publication and row creation) + is INSERTed with full identity rather than no-opped, so a gateway row is never + first-created by the identity-less lazy writer (``update_token_counts``) and left + unroutable forever.""" if not session_id or not session_key: return 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 = [*identity, session_id] + ancestors = include_compression_ancestors + query_params = [session_id, *identity] if ancestors else [*identity, session_id] def _do(conn): conn.execute( - f"""{lineage_cte} + f"""{_COMPRESSION_LINEAGE_CTE if ancestors else ""} UPDATE sessions SET session_key = ?, source = ?, user_id = ?, chat_id = ?, chat_type = ?, thread_id = ?, display_name = COALESCE(?, display_name), origin_json = COALESCE(?, origin_json) - {target_clause}""", + {"WHERE id IN (SELECT id FROM compression_lineage)" if ancestors else "WHERE id = ?"}""", query_params, ) - # Self-heal: the UPDATE silently no-ops on a missing row — insert it - # with full identity so the session is durably routable. - if include_compression_ancestors: + if ancestors: return - cur = conn.execute("SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,)) - if cur.fetchone() is None: + # The UPDATE silently no-ops on a missing row — insert it with full identity. + if conn.execute("SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,)).fetchone() is None: conn.execute( """INSERT INTO sessions ( id, source, user_id, session_key, chat_id, @@ -267,24 +244,17 @@ 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(), - ), + # Same ownership stamp as _insert_session_row: an unowned (NULL) row + # vanishes from profile-keyed consumers. + (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: - """Upsert one gateway routing entry (session_key -> SessionEntry JSON). - - ``gateway_routing`` durably replaces sessions.json. ``scope`` namespaces - the index per sessions_dir so two stores never share routing state. - """ + """Upsert one gateway routing entry (session_key -> SessionEntry JSON); ``scope`` + namespaces the index per sessions_dir so two stores never share routing state.""" if not session_key or not entry_json: return self._write_sql( @@ -297,11 +267,8 @@ class SessionGatewayMixin: ) 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 - write transaction; other scopes untouched. - """ + """Atomically replace the routing index for *scope* (keys absent from *entries* + are removed); other scopes untouched.""" now = time.time() def _do(conn): @@ -310,29 +277,21 @@ class SessionGatewayMixin: conn.executemany( "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], - ) + [(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]: """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,) - ) + rows = self._read_all("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]]: - """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, - tokens, tool/API calls, activity, or title): a leaked test fixture or a - chat routed but never answered. Safe to drop — no transcript to lose, - and the gateway mints a fresh session on the next inbound message. + """Keyed, still-open rows with no evidence of a single turn (no messages, tokens, + tool/API calls, activity, or title): leaked fixtures or chats routed but never + answered. Safe to drop — the gateway mints a fresh session on the next message. Needs its own selector because ``bulk prune``/``archive`` are pinned to - ``ended_at IS NOT NULL`` (never pick a live session), which excludes - every never-closed row. ``pinned``/``archived`` = explicit keep intent. - """ + ``ended_at IS NOT NULL``. ``pinned``/``archived`` = explicit keep intent.""" cutoff = time.time() - (float(older_than_days) * 86400.0) rows = self._read_all( """ @@ -362,16 +321,12 @@ class SessionGatewayMixin: return [dict(r) for r in rows] def _delete_routing_entries_for_sessions(self, session_ids: Set[str]) -> int: - """Drop ``gateway_routing`` rows pointing at any of *session_ids*. - - The target id lives only inside ``entry_json``, so matching is done in - Python over all scopes. - """ + """Drop ``gateway_routing`` rows pointing at any of *session_ids*; the target id + lives only inside ``entry_json``, so matching is done in Python over all scopes.""" if not session_ids: return 0 - rows = self._read_all("SELECT scope, session_key, entry_json FROM gateway_routing") doomed: List[Tuple[str, str]] = [] - for row in rows: + for row in self._read_all("SELECT scope, session_key, entry_json FROM gateway_routing"): try: entry = json.loads(row["entry_json"] or "{}") except Exception: @@ -380,22 +335,15 @@ 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 - ) + self._write_sql("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 - ) -> Tuple[int, int]: - """Delete never-active keyed rows and the routing entries naming them. - - Returns ``(sessions_deleted, routing_entries_deleted)``. Routing - entries go first: a stale entry outliving its target would have the - gateway resume a nonexistent session id. Deletion goes through - :meth:`delete_session` so the delegate cascade, FTS bookkeeping and - transcript cleanup stay owned by one implementation. - """ + self, *, older_than_days: float, sessions_dir: Optional[Path] = None) -> Tuple[int, int]: + """Delete never-active keyed rows and the routing entries naming them; returns + ``(sessions_deleted, routing_entries_deleted)``. Routing entries go first: a stale + entry outliving its target would have the gateway resume a nonexistent id. + Deletion goes through :meth:`delete_session` (delegate cascade, FTS, transcripts).""" candidates = self.list_never_active_keyed_sessions(older_than_days=older_than_days) if not candidates: return (0, 0) @@ -405,12 +353,10 @@ class SessionGatewayMixin: return (deleted, routing_deleted) def list_gateway_sessions( - 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``.""" - # Full rows carry token/cost totals — drain queued async accounting - # deltas so consumers see exact counters. + 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``.""" + # Full rows carry token/cost totals — drain queued async accounting deltas first. self.flush_token_counts() query = f""" SELECT sessions.*, @@ -426,44 +372,30 @@ class SessionGatewayMixin: WHERE s2.session_key = sessions.session_key ) """ - params: list = [] - if platform: - query += " AND LOWER(source) = LOWER(?)" - params.append(platform) - if active_only: - query += " AND ended_at IS NULL" - query += " ORDER BY last_active DESC" + params: list = [platform] if platform else [] + query += (" AND LOWER(source) = LOWER(?)" if platform else "") + ( + " AND ended_at IS NULL" if active_only else "") + " ORDER BY last_active DESC" 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, - thread_id: 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. - ``sessions.json`` is the fast index but can be missing or pruned; the - durable ``session_key`` on the row rebuilds the mapping exactly. Rows - ended only by the old ``agent_close`` bug or a mistaken TUI - ``ws_orphan_reap`` are recoverable; explicit boundaries (/new, /resume - switches, compression splits) are not. - - Ranked by ``last_activity_at`` (falling back to ``started_at``) — - ``started_at`` alone resurrected days-old zombie rows. Rows with - messages win, but an empty keyed row is still returned rather than - ``None`` (``None`` mints a brand-new session id; the transcript may - live under a compression child). Reset fence: a candidate is rejected - when a peer boundary row (``session_reset`` or any non-recoverable - end_reason) ended *after* its last activity, or the has-messages - ranking could reach behind a /new and restore the reset context. - - Fallback for a temporarily-missing exact key still requires the - complete peer tuple (never cross chats/threads/users) plus 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. - """ + The durable ``session_key`` on the row rebuilds a missing/pruned ``sessions.json`` + mapping. Rows ended only by the old ``agent_close`` bug or a mistaken TUI + ``ws_orphan_reap`` are recoverable; explicit boundaries (/new, /resume switches, + compression splits) are not. Ranked by ``last_activity_at`` (fallback + ``started_at`` — alone it resurrected days-old zombies); rows with messages win, + but an empty keyed row still beats ``None`` (which mints a new id while the + transcript may live under a compression child). Reset fence: a candidate is + rejected when a peer boundary row ended *after* its last activity, or the + has-messages ranking could reach behind a /new. The exact-key fallback requires + the complete peer tuple (never cross chats/threads/users) plus a profile fence: a + Telegram DM's tuple is identical for every bot, so a sibling profile's legacy row + would otherwise be adopted. Ours = 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: @@ -474,27 +406,19 @@ class SessionGatewayMixin: return None owner = self._own_profile_name() row = conn.execute( - _PEER_BY_TUPLE_SQL, - (source, user_id, chat_id, chat_type, thread_id, owner, owner, owner), + _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]]: - """Report message-bearing session rows that lost their routing identity. - - A candidate orphan has messages but no ``session_key``; it is - *adoptable* only when exactly one keyed predecessor can be named: - - * ``lineage`` — ``parent_session_id`` points at a keyed row of the - same source (a recorded fact; no time window). - * ``contiguity`` — exactly one keyed row of the same source (and - compatible ``user_id``) fell quiet within *max_gap_s* of the - orphan's start, and is older than the orphan's own last activity. - - Ambiguity is reported ``adoptable=False`` with a reason, never guessed: - mis-adopting splices one person's conversation into another's chat. - Branch/delegate/tool rows are excluded — unkeyed by design, not damage. - """ + """Report message-bearing rows that lost their routing identity (messages, no + ``session_key``). Adoptable only when exactly one keyed predecessor can be named: + ``lineage`` (``parent_session_id`` is a keyed row of the same source; no time + window) or ``contiguity`` (exactly one keyed same-source row with compatible + ``user_id`` fell quiet within *max_gap_s* of the orphan's start and is older than + its last activity). Ambiguity is reported ``adoptable=False`` with a reason, never + guessed — 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) records: List[Dict[str, Any]] = [] with self._read_ctx() as conn: @@ -504,8 +428,7 @@ class SessionGatewayMixin: if orphan["parent_session_id"]: evidence = "lineage" donor = conn.execute( - _ORPHAN_LINEAGE_DONOR_SQL, (orphan["parent_session_id"], orphan["source"]) - ).fetchone() + _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" else: @@ -525,15 +448,10 @@ class SessionGatewayMixin: 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, + "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. + "evidence": evidence if donor else "", "adoptable": donor is not None, "reason": reason}) + # Two unkeyed successors claiming one predecessor: at most one continues that chat. 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 @@ -546,11 +464,8 @@ class SessionGatewayMixin: 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 - gateway that healed either row makes this a no-op, not a conflicting - write. Non-NULL orphan columns are preserved. True when applied. - """ + Re-verifies the pair inside the write txn so a concurrent gateway that healed + 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 @@ -561,13 +476,9 @@ class SessionGatewayMixin: (donor_id,), ).fetchone() orphan = conn.execute( - "SELECT session_key, source FROM sessions WHERE id = ?", (orphan_id,) - ).fetchone() - if donor is None or orphan is None: - return False - if not donor["session_key"] or orphan["session_key"]: - return False - if (donor["source"] or "") != (orphan["source"] or ""): + "SELECT session_key, source FROM sessions WHERE id = ?", (orphan_id,)).fetchone() + if (donor is None or orphan is None or not donor["session_key"] or orphan["session_key"] + or (donor["source"] or "") != (orphan["source"] or "")): return False conn.execute( """UPDATE sessions @@ -583,14 +494,12 @@ class SessionGatewayMixin: (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 - # running and the newly keyed orphan could lose the chat again. + # Retire under a reason recovery does NOT treat as resumable — 'agent_close' / + # 'ws_orphan_reap' would keep it in the running and the orphan could lose the chat again. conn.execute( "UPDATE sessions SET ended_at = COALESCE(ended_at, ?), " "end_reason = 'superseded_by_repair' WHERE id = ?", - (time.time(), donor_id), - ) + (time.time(), donor_id)) return True return self._execute_write(_do) @@ -609,8 +518,7 @@ class SessionGatewayMixin: (session_key,), ) row = conn.execute( - "SELECT failure_streak FROM gateway_hygiene_state WHERE session_key = ?", - (session_key,), + "SELECT failure_streak FROM gateway_hygiene_state WHERE session_key = ?", (session_key,), ).fetchone() return int(row[0]) @@ -624,15 +532,11 @@ class SessionGatewayMixin: @staticmethod def session_gateway_runtime(session_meta: Optional[Dict[str, Any]]) -> Dict[str, Any]: - """Read the persisted runtime route off a session row dict. - - Accepts ``get_session``'s dict (``model_config`` as JSON string) or a - parsed dict. Precedence: nested ``gateway_runtime`` (gateway sync / CLI - ``/model``), then top-level ``provider``/``base_url``/``api_mode`` (TUI - ``_runtime_model_config``), then ``billing_provider`` so sessions that - never ran ``/model`` still restore the provider that served them. - Empty dict on parse failure — resume uses ambient config. - """ + """Read the persisted runtime route off a session row dict (``model_config`` as + JSON string or parsed dict). Precedence: nested ``gateway_runtime`` (gateway sync / + CLI ``/model``), then top-level ``provider``/``base_url``/``api_mode`` (TUI), then + ``billing_provider`` so sessions that never ran ``/model`` still restore the + provider that served them. Empty dict on parse failure — resume uses ambient config.""" from hermes_state import _BARE_BILLING_PROVIDERS raw = (session_meta or {}).get("model_config") if isinstance(raw, str): @@ -643,93 +547,74 @@ 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. + # 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 survives here. if isinstance(runtime, dict) and runtime.get("provider"): 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)} 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 is 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 default. 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 {} + if not isinstance(runtime, dict): + return {} + return {k: v for k, v in runtime.items() if v is not None} def register_backend_heartbeat( - 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. - - ``backend_id`` MUST be stable for the process lifetime (e.g. - ``f"{profile}@{host}:{pid}"``) so a respawn cannot inherit a dead - predecessor's heartbeat and protect stale rows. ``started_at`` is when - THIS process started, not first-refresh wall clock, so a backend whose - previous run died is not mistaken for a freshly-spawned sibling. - """ + 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. ``backend_id`` MUST be stable for the process + lifetime (e.g. ``f"{profile}@{host}:{pid}"``) so a respawn cannot inherit a dead + predecessor's heartbeat; ``started_at`` is when THIS process started, so a backend + whose previous run died is not mistaken for a freshly-spawned sibling.""" if not backend_id: return ts = time.time() if last_heartbeat is None else float(last_heartbeat) self._write_sql( - "INSERT INTO gateway_heartbeats" - " (backend_id, pid, started_at, last_heartbeat, profile, host)" - " VALUES (?, ?, ?, ?, ?, ?)" - " ON CONFLICT(backend_id) DO UPDATE SET" - " pid = excluded.pid," - " started_at = excluded.started_at," - " last_heartbeat = excluded.last_heartbeat," - " profile = excluded.profile," - " host = excluded.host", - (str(backend_id), int(pid), float(started_at), ts, str(profile), str(host)), - ) + "INSERT INTO gateway_heartbeats (backend_id, pid, started_at, last_heartbeat, profile, host)" + " VALUES (?, ?, ?, ?, ?, ?) ON CONFLICT(backend_id) DO UPDATE SET pid = excluded.pid," + " started_at = excluded.started_at, last_heartbeat = excluded.last_heartbeat," + " profile = excluded.profile, host = excluded.host", + (str(backend_id), int(pid), float(started_at), ts, str(profile), str(host))) def clear_backend_heartbeat(self, backend_id: str) -> bool: - """Remove this backend's heartbeat row (from ``atexit``); True if removed. - A crashed backend's row is reclaimed later by ``prune_stale_heartbeats``.""" + """Remove this backend's heartbeat row (from ``atexit``); True if removed. A crashed + backend's row is reclaimed later by ``prune_stale_heartbeats``.""" if not backend_id: return False return self._write_rowcount( - "DELETE FROM gateway_heartbeats WHERE backend_id = ?", (str(backend_id),) - ) > 0 + "DELETE FROM gateway_heartbeats WHERE backend_id = ?", (str(backend_id),)) > 0 def prune_stale_heartbeats(self, *, max_age_seconds: float) -> List[str]: - """Drop heartbeat rows older than the staleness window; return removed - backend ids. Safe from any process — only stale rows are touched.""" + """Drop heartbeat rows older than the staleness window; return removed backend ids. + Safe from any process — only stale rows are touched.""" 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", - (cutoff,), - ) + "DELETE FROM gateway_heartbeats WHERE last_heartbeat < ? RETURNING backend_id", + (cutoff,)) return [str(r[0]) for r in cur.fetchall()] return list(self._execute_write(_do) or []) def list_backend_heartbeats(self) -> List[Dict[str, Any]]: """Snapshot of every backend heartbeat (diagnostics/tests); fields mirror the table.""" rows = self._read_all( - "SELECT backend_id, pid, started_at, last_heartbeat," - " profile, host FROM gateway_heartbeats" - " ORDER BY last_heartbeat DESC", - ) + "SELECT backend_id, pid, started_at, last_heartbeat, profile, host FROM gateway_heartbeats" + " ORDER BY last_heartbeat DESC") 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.""" return self._write_rowcount( - "UPDATE sessions " - "SET handoff_state = 'pending', " - " handoff_platform = ?, " - " handoff_error = NULL " - "WHERE id = ? AND (handoff_state IS NULL " + "UPDATE sessions SET handoff_state = 'pending', handoff_platform = ?, " + " handoff_error = NULL WHERE id = ? AND (handoff_state IS NULL " " OR handoff_state IN ('completed', 'failed'))", (platform, session_id), ) > 0 @@ -738,17 +623,12 @@ class SessionGatewayMixin: """Return ``{"state", "platform", "error"}`` or None if the session has no handoff record.""" try: row = self._read_one( - "SELECT handoff_state, handoff_platform, handoff_error " - "FROM sessions WHERE id = ?", - (session_id,), - ) + "SELECT handoff_state, handoff_platform, handoff_error FROM sessions WHERE id = ?", + (session_id,)) if not row: return None - return { - "state": row["handoff_state"], - "platform": row["handoff_platform"], - "error": row["handoff_error"], - } + return {"state": row["handoff_state"], "platform": row["handoff_platform"], + "error": row["handoff_error"]} except Exception: return None @@ -756,13 +636,10 @@ class SessionGatewayMixin: """All sessions in handoff_state='pending', oldest first (gateway handoff watcher).""" try: rows = self._read_all( - "SELECT s.*, " - "COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " - "FROM sessions s " + "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", - ) + "ORDER BY s.started_at ASC") return [self._session_row_dict(r) for r in rows] except Exception: return [] @@ -770,77 +647,54 @@ class SessionGatewayMixin: def claim_handoff(self, session_id: str) -> bool: """Atomically transition pending → running. Returns True if claimed.""" return self._write_rowcount( - "UPDATE sessions SET handoff_state = 'running' " - "WHERE id = ? AND handoff_state = 'pending'", + "UPDATE sessions SET handoff_state = 'running' WHERE id = ? AND handoff_state = 'pending'", (session_id,), ) > 0 def complete_handoff(self, session_id: str) -> None: """Mark a handoff as completed.""" self._write_sql( - "UPDATE sessions SET handoff_state = 'completed', " - "handoff_error = NULL WHERE id = ?", - (session_id,), - ) + "UPDATE sessions SET handoff_state = 'completed', handoff_error = NULL WHERE id = ?", + (session_id,)) def fail_handoff( - self, session_id: str, error: str, *, only_states: Optional[Tuple[str, ...]] = None - ) -> bool: + 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. - ``only_states`` makes the write a compare-and-swap on ``handoff_state``. - Waiters that give up (CLI 60s poll, Desktop bounded poll) MUST pass - ``only_states=("pending",)``: once the gateway watcher has claimed the - row (``running``) it owns the terminal state, and an unconditional - waiter-side fail races the dispatch — the gateway later overwrites - ``failed`` → ``completed`` after the user was told the gateway is down - (split-brain: the handoff delivered and ``switch_session`` re-pointed - the session). The watcher fails its OWN claimed row unconditionally. - """ - 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 + ``only_states`` makes the write a compare-and-swap on ``handoff_state``. Waiters + that give up (CLI 60s poll, Desktop bounded poll) MUST pass ``only_states=("pending",)``: + once the watcher has claimed the row (``running``) it owns the terminal state, and an + unconditional waiter-side fail races the dispatch — the gateway later overwrites + ``failed`` → ``completed`` after the user was told the gateway is down (split-brain: + the handoff delivered and ``switch_session`` re-pointed the session). The watcher + fails its OWN claimed row unconditionally.""" + states = tuple(only_states) if only_states else () + sql = _HANDOFF_FAIL_SQL + "id = ?" + ( + f" AND handoff_state IN ({', '.join('?' for _ in states)})" if states else "") + return self._write_rowcount(sql, (error[:500], session_id, *states)) > 0 def reclaim_stale_running_handoffs(self, error: str) -> List[str]: - """Fail every handoff stuck in ``running``. Returns the ids reclaimed. + """Fail every handoff stuck in ``running``; returns the ids reclaimed. - Only the gateway watcher sets ``running``, and only for one in-process - dispatch — so a ``running`` row at watcher startup belongs to a PREVIOUS - gateway that died mid-dispatch. It is poisonous: ``request_handoff`` - only accepts NULL/``completed``/``failed``, so the session could never - hand off again, with no error surfaced. Failing rather than re-queueing - is deliberate: the dead gateway may already have switched the session - key and dispatched the synthetic turn, so a blind retry risks double - delivery; a clean terminal state the user can retry from is right. - """ + Only the gateway watcher sets ``running``, for one in-process dispatch — so a + ``running`` row at watcher startup belongs to a PREVIOUS gateway that died + mid-dispatch. It is poisonous: ``request_handoff`` only accepts NULL/``completed``/ + ``failed``, so the session could never hand off again, with no error surfaced. + Failing rather than re-queueing is deliberate: the dead gateway may already have + switched the session key and dispatched the synthetic turn, so a blind retry risks + double 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'") ids = [r[0] for r in cur.fetchall()] if ids: - conn.execute( - "UPDATE sessions SET handoff_state = 'failed', " - "handoff_error = ? WHERE handoff_state = 'running'", - (error[:500],), - ) + conn.execute(_HANDOFF_FAIL_SQL + "handoff_state = 'running'", (error[:500],)) return ids try: return self._execute_write(_do) or [] except Exception: - # Swallow but never silently: a persistently failing reclaim leaves - # poisonous 'running' rows in place, so the operator needs a trace. + # Swallow but never silently: a persistently failing reclaim leaves poisonous + # 'running' rows in place, so the operator needs a trace. logger.warning( "reclaim_stale_running_handoffs failed; stranded 'running' " - "handoff rows (if any) were left in place", exc_info=True, - ) + "handoff rows (if any) were left in place", exc_info=True) return []