diff --git a/hermes_state_portability.py b/hermes_state_portability.py index f31016d16e..1d23308a46 100644 --- a/hermes_state_portability.py +++ b/hermes_state_portability.py @@ -102,9 +102,6 @@ class SessionPortabilityMixin: with self._lock: return self._conn.execute(sql, params).fetchall() - def _rich_rows(self, sql: str, params=()) -> List[Dict[str, Any]]: - return [self._rich_row(row) for row in self._locked_rows(sql, params)] - def distinct_session_cwds(self, include_archived: bool = False) -> List[Dict[str, Any]]: """Distinct non-empty session cwds with usage stats, for repo discovery. Aggregates across ALL history; children/branches count (a worktree session is a real @@ -116,10 +113,8 @@ class SessionPortabilityMixin: "SELECT cwd AS cwd, COUNT(*) AS sessions, MAX(COALESCE(ended_at, started_at, 0)) AS last_active " f"FROM sessions WHERE {where} GROUP BY cwd" ) - return [ - {"cwd": r["cwd"], "sessions": int(r["sessions"] or 0), "last_active": float(r["last_active"] or 0)} - for r in rows - ] + return [{"cwd": r["cwd"], "sessions": int(r["sessions"] or 0), "last_active": float(r["last_active"] or 0)} + for r in rows] def list_cron_job_runs(self, job_id: str, limit: int = 20, offset: int = 0) -> List[Dict[str, Any]]: """Run sessions of one cron job, newest first, in the ``list_sessions_rich`` row shape. @@ -135,7 +130,7 @@ class SessionPortabilityMixin: "\n ORDER BY s.started_at DESC, s.id DESC\n LIMIT ? OFFSET ?", prompt_select=f",\n {_PROMPT_RESOLVED_SQL}", ) - return self._rich_rows(query, (prefix, prefix_hi, limit, offset)) + return [self._rich_row(row) for row in self._locked_rows(query, (prefix, prefix_hi, limit, offset))] def _get_session_rich_row(self, session_id: str, compact_rows: bool = False) -> Optional[Dict[str, Any]]: """One session with the ``list_sessions_rich`` enriched columns, or None. @@ -165,14 +160,13 @@ class SessionPortabilityMixin: self._compact_session_cols() if compact_rows else "s.*", f"s.id IN ({','.join('?' for _ in ids)})", prompt_select=None if compact_rows else f", {_PROMPT_RESOLVED_SQL}", ) - return {s["id"]: s for s in self._rich_rows(query, ids)} + return {s["id"]: s for s in map(self._rich_row, self._locked_rows(query, ids))} def list_skill_scaffolded_sessions(self, limit: int = 200) -> List[Dict[str, Any]]: """Titled sessions whose first user turn was a ``/skill`` invocation (their titles describe the expanded skill body, not the request). Returns ``id``, ``title`` and the first-turn ``content`` so callers can re-derive what was typed. Newest first.""" - rows = self._locked_rows( - """ + rows = self._locked_rows(""" SELECT s.id, s.title, m.content FROM sessions s JOIN messages m ON m.id = ( @@ -184,9 +178,7 @@ class SessionPortabilityMixin: WHERE s.title IS NOT NULL AND m.content LIKE ? ORDER BY s.started_at DESC LIMIT ? - """, - (SKILL_SCAFFOLD_SQL_LIKE, int(limit)), - ) + """, (SKILL_SCAFFOLD_SQL_LIKE, int(limit))) return [dict(row) for row in rows] # ── Export ───────────────────────────────────────────────────────────── @@ -230,10 +222,8 @@ class SessionPortabilityMixin: ``adopted`` and ``donor_retired`` (True only when EVERY segment retired).""" payload = donor_db.export_session_lineage(session_id) if not payload: - return { - "ok": False, "adopted": False, "donor_retired": False, - "error": f"session {session_id!r} not found in donor store", - } + return {"ok": False, "adopted": False, "donor_retired": False, + "error": f"session {session_id!r} not found in donor store"} segments = payload.get("segments") or [payload] # Divergence guard: a segment we will SKIP (already here) may have kept growing in @@ -248,21 +238,16 @@ class SessionPortabilityMixin: local_count = len(self.get_messages(seg_id)) if donor_count > local_count: donor_ahead = True - logger.warning( - "adoption divergence: donor segment %s has %d messages, " - "local copy has %d — donor will NOT be retired", - seg_id, donor_count, local_count, - ) + logger.warning("adoption divergence: donor segment %s has %d messages, " + "local copy has %d — donor will NOT be retired", seg_id, donor_count, local_count) result = self.import_sessions([dict(seg) for seg in segments]) imported = int(result.get("imported") or 0) skipped = int(result.get("skipped") or 0) adopted = result.get("ok", False) and (imported + skipped) == len(segments) if not adopted: - logger.warning( - "adoption of %s did not complete: imported=%s skipped=%s of %s segment(s); errors=%s", - session_id, imported, skipped, len(segments), result.get("errors"), - ) + logger.warning("adoption of %s did not complete: imported=%s skipped=%s of %s segment(s); errors=%s", + session_id, imported, skipped, len(segments), result.get("errors")) donor_retired = False if adopted and retire_donor and not donor_ahead: @@ -282,8 +267,7 @@ class SessionPortabilityMixin: if donor_now > local_now: logger.warning( "adoption divergence at retire time: donor segment %s grew to %d messages (local %d) — " - "leaving donor unretired", - seg_id, donor_now, local_now, + "leaving donor unretired", seg_id, donor_now, local_now, ) return False # First end_reason wins in end_session(); reopen so the adoption boundary is @@ -338,26 +322,12 @@ class SessionPortabilityMixin: except (TypeError, ValueError): return default - @staticmethod - def _reasoning_json_value(value: Any) -> Any: - return safe_json_loads(value, default=value) if isinstance(value, str) else value - - @staticmethod - def _import_error(index: int, session_id: str, error: str) -> Dict[str, Any]: - item: Dict[str, Any] = {"index": index, "error": error} - if session_id: - item["session_id"] = session_id - return item - def _normalize_import_session(self, raw: Dict[str, Any], session_id: str, messages: list) -> Dict[str, Any]: """Type-check one payload session + its messages; raises ValueError.""" clean_session = dict(raw) clean_session["id"] = session_id clean_session["model_config"] = self._import_json_object_or_none(clean_session.get("model_config"), "model_config") - clean_session["parent_session_id"] = self._import_text_or_none( - clean_session.get("parent_session_id"), "parent_session_id" - ) - for field in _IMPORT_SESSION_TEXT_FIELDS: + for field in ("parent_session_id", *_IMPORT_SESSION_TEXT_FIELDS): clean_session[field] = self._import_text_or_none(clean_session.get(field), field) clean_messages: List[Dict[str, Any]] = [] for message_index, message in enumerate(messages): @@ -383,7 +353,10 @@ class SessionPortabilityMixin: try: item = self._validate_import_session(raw, session_id, seen_ids, totals) except ValueError as exc: - errors.append(self._import_error(index, session_id, str(exc))) + item = {"index": index, "error": str(exc)} + if session_id: + item["session_id"] = session_id + errors.append(item) continue seen_ids.add(session_id) normalized.append({"index": index, **item}) @@ -433,15 +406,14 @@ class SessionPortabilityMixin: **{col: self._coerce_or(raw.get(col), int, 0) for col in _IMPORT_INT_COLS}, } conn.execute(_IMPORT_SESSION_INSERT_SQL, params) + def _json_value(value: Any) -> Any: + return safe_json_loads(value, default=value) if isinstance(value, str) else value sanitized_messages = [ - {**msg, **{key: self._reasoning_json_value(msg.get(key)) for key in _IMPORT_MESSAGE_JSON_FIELDS}} - for msg in messages + {**msg, **{key: _json_value(msg.get(key)) for key in _IMPORT_MESSAGE_JSON_FIELDS}} for msg in messages ] total_messages, total_tool_calls = self._insert_message_rows(conn, session_id, sanitized_messages) - conn.execute( - "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", - (total_messages, total_tool_calls, session_id), - ) + conn.execute("UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", + (total_messages, total_tool_calls, session_id)) @staticmethod def _attach_import_parents(conn, parent_updates: List[tuple]) -> int: @@ -509,11 +481,10 @@ class SessionPortabilityMixin: if parent_id: parent_updates.append((session_id, parent_id)) imported_ids.append(session_id) - detached = self._attach_import_parents(conn, parent_updates) return { "ok": True, "imported": len(imported_ids), "skipped": len(skipped_ids), - "detached": detached, "imported_ids": imported_ids, "skipped_ids": skipped_ids, - "errors": [], + "detached": self._attach_import_parents(conn, parent_updates), + "imported_ids": imported_ids, "skipped_ids": skipped_ids, "errors": [], } return self._execute_write(_do) diff --git a/hermes_state_registry.py b/hermes_state_registry.py index cb6276d1dc..51cc6ef8cf 100644 --- a/hermes_state_registry.py +++ b/hermes_state_registry.py @@ -20,6 +20,7 @@ Lifecycle rules: from __future__ import annotations +import contextlib import logging import threading from pathlib import Path @@ -57,7 +58,7 @@ _opening: Dict[Path, threading.Event] = {} def _open_session_db(path: Path) -> "SessionDB": - """Construct the SessionDB for *path* (call-time import avoids cycles).""" + """Construct the SessionDB for *path* (call-time import avoids cycles; tests patch this).""" from hermes_state import SessionDB return SessionDB(db_path=path) @@ -65,10 +66,8 @@ def _open_session_db(path: Path) -> "SessionDB": def _teardown(db: "SessionDB") -> None: """Close a shared instance, clearing its registry-owned flag first.""" - try: + with contextlib.suppress(Exception): db._shared_registry_owned = False - except Exception: - pass try: db.close() except Exception: @@ -78,10 +77,8 @@ def _teardown(db: "SessionDB") -> None: def _db_path_of(db: "SessionDB") -> Optional[Path]: """``Path(db.db_path)`` or None when absent/unconvertible.""" path = getattr(db, "db_path", None) - if path is None: - return None try: - return Path(path) + return None if path is None else Path(path) except (TypeError, ValueError): return None @@ -113,8 +110,12 @@ def acquire(db_path: Optional[Path] = None) -> "SessionDB": if generation is not None: current = _stat_db_file_identity(path) if current is not None and generation.identity is not None and current != generation.identity: - # File replaced: retire, then elect one caller to open the replacement. - _retire_generation_locked(path, generation) + # File replaced: retire this generation so it is never lent again, then elect one + # caller to open the replacement. It stays alive for its holders, tracked in + # ``_retired`` by ``id(db)`` so their releases find it after the path remaps. + generation.retired = True + del _generations[path] + _retired[id(generation.db)] = generation else: generation.refcount += 1 return generation.db @@ -139,8 +140,7 @@ def acquire(db_path: Optional[Path] = None) -> "SessionDB": with _lock: existing = _generations.get(path) - if existing is not None: - # Defensive: installed by explicit registry manipulation mid-open. + if existing is not None: # Defensive: installed by explicit registry manipulation mid-open. existing.refcount += 1 winner = existing.db else: @@ -152,16 +152,6 @@ def acquire(db_path: Optional[Path] = None) -> "SessionDB": return winner -def _retire_generation_locked(path: Path, generation: _Generation) -> None: - """Retire *generation* so it is never lent again (caller holds _lock). It stays alive - for its holders, tracked in ``_retired`` by ``id(db)`` so their releases find it even - after the path maps to a new generation.""" - generation.retired = True - if _generations.get(path) is generation: - del _generations[path] - _retired[id(generation.db)] = generation - - def release(db: "SessionDB") -> bool: """Decrement the refcount of a shared SessionDB. ``True`` if *db* was shared; ``False`` if it is not registry-managed (caller owns close()). The final release tears the @@ -182,13 +172,10 @@ def release(db: "SessionDB") -> bool: return False generation.refcount -= 1 needs_teardown = generation.refcount <= 0 - if needs_teardown: - if generation.retired: - _retired.pop(key, None) - else: - path = _db_path_of(db) - if path is not None: - _generations.pop(path, None) + if needs_teardown and generation.retired: + _retired.pop(key, None) + elif needs_teardown and (path := _db_path_of(db)) is not None: + _generations.pop(path, None) # Teardown OUTSIDE the lock: stopping the token writer, WAL checkpoint and read-pool # drain must not block acquisition for every other state.db. if needs_teardown: @@ -222,8 +209,7 @@ def stats() -> Dict[str, int]: """Registry census for tests and diagnostics (no locks held long).""" with _lock: return { - "live_generations": len(_generations), - "retired_generations": len(_retired), + "live_generations": len(_generations), "retired_generations": len(_retired), "total_refcounts": sum(g.refcount for g in _generations.values()), } diff --git a/hermes_state_schema.py b/hermes_state_schema.py index 7945013d43..d4387cf239 100644 --- a/hermes_state_schema.py +++ b/hermes_state_schema.py @@ -4,6 +4,7 @@ Plain mixin for ``hermes_state.SessionDB`` (no ``__init__``/state of its own). Must never import hermes_state (cycle); shared constants live in hermes_state_common. """ +import contextlib import datetime import hashlib import logging @@ -196,10 +197,8 @@ class SessionSchemaMixin: def _drop_all_fts_triggers(self, cursor: sqlite3.Cursor) -> None: self._drop_fts_triggers(cursor) for trigger in _FTS_CJK_TRIGGERS: - try: + with contextlib.suppress(sqlite3.OperationalError): cursor.execute(f"DROP TRIGGER IF EXISTS {trigger}") - except sqlite3.OperationalError: - pass @staticmethod def _fts_triggers_missing(cursor: sqlite3.Cursor, names: Sequence[str]) -> bool: @@ -207,10 +206,8 @@ class SessionSchemaMixin: if not names: return False # "name IN ()" is a SQLite syntax error placeholders = ",".join("?" for _ in names) - row = cursor.execute( - f"SELECT COUNT(*) FROM sqlite_master WHERE type = 'trigger' AND name IN ({placeholders})", tuple(names), - ).fetchone() - return int(row[0]) < len(names) + sql = f"SELECT COUNT(*) FROM sqlite_master WHERE type = 'trigger' AND name IN ({placeholders})" + return int(cursor.execute(sql, tuple(names)).fetchone()[0]) < len(names) @staticmethod def _fts_update_trigger_needs_narrowing(sql: Optional[str]) -> bool: @@ -227,13 +224,12 @@ class SessionSchemaMixin: # CJK is v23-only. Decide the layout before selecting destructive candidates so the # legacy branch never drops a trigger it won't recreate. legacy_layout = self._db_has_legacy_inline_fts(cursor) - update_names = ("messages_fts_update", "messages_fts_trigram_update") - if not legacy_layout: - update_names += ("messages_fts_cjk_update",) + update_names = ("messages_fts_update", "messages_fts_trigram_update") + ( + () if legacy_layout else ("messages_fts_cjk_update",) + ) placeholders = ", ".join("?" for _ in update_names) - rows = cursor.execute( - f"SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND name IN ({placeholders})", update_names, - ).fetchall() + sql = f"SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND name IN ({placeholders})" + rows = cursor.execute(sql, update_names).fetchall() to_drop = [name for name, sql in rows if self._fts_update_trigger_needs_narrowing(sql)] if not to_drop: return 0 @@ -253,7 +249,10 @@ class SessionSchemaMixin: self._quarantine_cjk_after_update_of_migration(cursor) logger.exception("CJK FTS re-ensure after UPDATE OF migration failed") raise - if not self._cjk_update_trigger_is_narrowed(cursor): + row = cursor.execute( + "SELECT sql FROM sqlite_master WHERE type = 'trigger' AND name = ?", ("messages_fts_cjk_update",), + ).fetchone() + if not row or self._fts_update_trigger_needs_narrowing(row[0]): self._quarantine_cjk_after_update_of_migration(cursor) logger.warning( "CJK FTS UPDATE trigger missing or still broad after " @@ -262,13 +261,6 @@ class SessionSchemaMixin: logger.info("Migrated %d broad FTS UPDATE trigger(s) to AFTER UPDATE OF (no rebuild required)", len(to_drop)) return len(to_drop) - def _cjk_update_trigger_is_narrowed(self, cursor: sqlite3.Cursor) -> bool: - """True when messages_fts_cjk_update exists with AFTER UPDATE OF.""" - row = cursor.execute( - "SELECT sql FROM sqlite_master WHERE type = 'trigger' AND name = ?", ("messages_fts_cjk_update",), - ).fetchone() - return bool(row) and not self._fts_update_trigger_needs_narrowing(row[0]) - def _quarantine_cjk_after_update_of_migration(self, cursor: sqlite3.Cursor) -> None: """Fail closed after dropping the CJK UPDATE trigger mid-migration: clear availability, persist ``fts_cjk_stale``, drop any residual trigger so a later open cannot @@ -284,22 +276,20 @@ class SessionSchemaMixin: logger.debug("Could not drop residual CJK UPDATE trigger after quarantine", exc_info=True) @staticmethod - def _rebuild_fts_indexes(cursor: sqlite3.Cursor, *, include_trigram: bool = True) -> None: + def _rebuild_fts_indexes(cursor: sqlite3.Cursor, *, legacy: bool = False, include_trigram: bool = True) -> None: """v23+ external-content 'rebuild'. It indexes EVERY row, so the deferred-backfill - markers are cleared or the worker would re-insert covered rows (duplicates).""" - cursor.execute("INSERT INTO messages_fts(messages_fts) VALUES('rebuild')") - if include_trigram: - cursor.execute("INSERT INTO messages_fts_trigram(messages_fts_trigram) VALUES('rebuild')") - cursor.execute(_CLEAR_REBUILD_MARKERS_SQL) - - @staticmethod - def _rebuild_legacy_fts_indexes(cursor: sqlite3.Cursor, *, include_trigram: bool = True) -> None: - """Rebuild the LEGACY inline (pre-v23) FTS indexes: no external-content 'rebuild' source, - so DELETE + reinsert the concatenated content the legacy triggers produced.""" + markers are cleared or the worker would re-insert covered rows (duplicates). + ``legacy`` (pre-v23 inline layout) has no external-content 'rebuild' source, so it + DELETEs + reinserts the concatenated content the legacy triggers produced.""" tables = ("messages_fts", "messages_fts_trigram") if include_trigram else ("messages_fts",) for tbl in tables: - cursor.execute(f"DELETE FROM {tbl}") - cursor.execute(f"INSERT INTO {tbl}(rowid, content) SELECT id, {_LEGACY_INLINE_CONCAT_SQL}FROM messages") + if legacy: + cursor.execute(f"DELETE FROM {tbl}") + cursor.execute(f"INSERT INTO {tbl}(rowid, content) SELECT id, {_LEGACY_INLINE_CONCAT_SQL}FROM messages") + else: + cursor.execute(f"INSERT INTO {tbl}({tbl}) VALUES('rebuild')") + if not legacy: + cursor.execute(_CLEAR_REBUILD_MARKERS_SQL) def _fts_table_probe(self, cursor: sqlite3.Cursor, table_name: str) -> Optional[bool]: """True = queryable, False = absent, None = FTS module/tokenizer missing or content @@ -328,9 +318,7 @@ class SessionSchemaMixin: decode_exc = exc logger.warning( "%s probe encountered invalid UTF-8 in FTS content; " - "search may return incomplete results until FTS is rebuilt: %s", - table_name, - decode_exc, + "search may return incomplete results until FTS is rebuilt: %s", table_name, decode_exc, ) return None @@ -342,17 +330,14 @@ class SessionSchemaMixin: ``_FTS_HOLDER_ESCALATE_SECONDS``, provably inactive orphan Desktop backends are reaped and the holders re-checked.""" now = time.time() - record = {} try: row = cursor.execute( "SELECT value FROM state_meta WHERE key = ? LIMIT 1", (FTS_REBUILD_DEFERRAL_KEY,), ).fetchone() except sqlite3.Error: row = None - if row: - parsed = safe_json_loads(row[0]) - if isinstance(parsed, dict): - record = parsed + parsed = safe_json_loads(row[0]) if row else None + record = parsed if isinstance(parsed, dict) else {} try: first_seen = float(record.get("first_seen", now)) attempts = int(record.get("attempts", 0)) + 1 @@ -375,9 +360,7 @@ class SessionSchemaMixin: if reaped: logger.error( "Reaped inactive orphan Desktop backend(s) %s after %d " - "state.db FTS rebuild deferrals; checking holders again.", - reaped, - attempts, + "state.db FTS rebuild deferrals; checking holders again.", reaped, attempts, ) foreign_holders = self._foreign_state_db_holders() if foreign_holders: @@ -385,17 +368,14 @@ class SessionSchemaMixin: "state.db FTS repair remains blocked after %d deferrals " "by holder(s) %s. Stop the listed processes, then run " "`hermes sessions optimize-storage` with the gateway stopped. " - "`hermes doctor` reports this degraded state.", - attempts, - foreign_holders, + "`hermes doctor` reports this degraded state.", attempts, foreign_holders, ) if not foreign_holders: return False logger.warning( "Deferred stale state.db FTS rebuild while foreign processes " "hold the database or WAL sidecars (%s); canonical writes and LIKE search remain available (deferral %d).", - foreign_holders, - attempts, + foreign_holders, attempts, ) return True @@ -409,8 +389,8 @@ class SessionSchemaMixin: with fts_rebuild_admission(self.db_path, timeout_seconds=timeout_seconds) as admitted: if not admitted: logger.warning( - "Deferred stale state.db FTS rebuild: another process " - "holds the rebuild authority; canonical writes and LIKE search remain available." + "Deferred stale state.db FTS rebuild: another process holds the rebuild authority; " + "canonical writes and LIKE search remain available." ) return False return self._recover_stale_fts_locked(cursor, legacy=legacy) @@ -445,10 +425,8 @@ class SessionSchemaMixin: # decides when it comes back online. self._ensure_fts_cjk_schema(cursor) self._fts_stale_retry_interval = 0.0 - try: + with contextlib.suppress(sqlite3.Error): self._conn.commit() - except sqlite3.Error: - pass return recovered except Exception: # noqa: BLE001 - background retry must never raise logger.warning( @@ -460,59 +438,43 @@ class SessionSchemaMixin: """Body of :meth:`_recover_stale_fts`; caller holds rebuild authority. One write transaction, so no canonical writer slips between rebuild and trigger restoration.""" try: - trigram_status = self._fts_table_probe(cursor, "messages_fts_trigram") + include_trigram = self._fts_table_probe(cursor, "messages_fts_trigram") is True except (sqlite3.DatabaseError, UnicodeDecodeError): # A corrupt vtable may fail even a LIMIT 0 probe; still include it in the drop-and-recreate. - trigram_status = True - include_trigram = trigram_status is True + include_trigram = True drop_sql = "".join(f"DROP TRIGGER IF EXISTS {trigger};" for trigger in _FTS_TRIGGERS) if include_trigram: drop_sql += "DROP TABLE IF EXISTS messages_fts_trigram;" - drop_sql += "DROP VIEW IF EXISTS messages_fts_trigram_src;" - drop_sql += "DROP TABLE IF EXISTS messages_fts;" - + drop_sql += "DROP VIEW IF EXISTS messages_fts_trigram_src;DROP TABLE IF EXISTS messages_fts;" if legacy: - schema_sql = LEGACY_FTS_SQL - if include_trigram: - schema_sql += LEGACY_FTS_TRIGRAM_SQL - rebuild_sql = schema_sql + _legacy_inline_reinsert_sql("messages_fts", 16) + rebuild_sql = LEGACY_FTS_SQL + (LEGACY_FTS_TRIGRAM_SQL if include_trigram else "") + rebuild_sql += _legacy_inline_reinsert_sql("messages_fts", 16) if include_trigram: rebuild_sql += _legacy_inline_reinsert_sql("messages_fts_trigram", 20, delete_first=True) else: - schema_sql = FTS_SQL - if include_trigram: - schema_sql += FTS_TRIGRAM_SQL - rebuild_sql = schema_sql + "INSERT INTO messages_fts(messages_fts) VALUES('rebuild');" + rebuild_sql = FTS_SQL + (FTS_TRIGRAM_SQL if include_trigram else "") + rebuild_sql += "INSERT INTO messages_fts(messages_fts) VALUES('rebuild');" if include_trigram: rebuild_sql += "INSERT INTO messages_fts_trigram(messages_fts_trigram) VALUES('rebuild');" rebuild_sql += _CLEAR_REBUILD_MARKERS_SQL + ";" - recovery_sql = ( - "BEGIN IMMEDIATE;" - + drop_sql - + rebuild_sql - + "DELETE FROM state_meta WHERE key IN " - + f"('{FTS_STALE_KEY}', '{FTS_REBUILD_DEFERRAL_KEY}');" - + "COMMIT;" + "BEGIN IMMEDIATE;" + drop_sql + rebuild_sql + + f"DELETE FROM state_meta WHERE key IN ('{FTS_STALE_KEY}', '{FTS_REBUILD_DEFERRAL_KEY}');COMMIT;" ) try: cursor.executescript(recovery_sql) except sqlite3.DatabaseError as exc: - try: + with contextlib.suppress(sqlite3.Error): self._conn.rollback() - except sqlite3.Error: - pass # Stale indexes must stay detached even on builds whose DDL transaction behavior differs. self._drop_all_fts_triggers(cursor) self._conn.commit() logger.error( "Automatic rebuild of stale FTS indexes failed (%s); " - "canonical writes remain enabled with FTS detached.", - exc, + "canonical writes remain enabled with FTS detached.", exc, ) return False - self._fts_stale = False self._fts_enabled = True self._trigram_available = include_trigram @@ -529,7 +491,7 @@ class SessionSchemaMixin: database still runs every startup. A corrupt/stale cache degrades to recomputation.""" cache_path = None schema_hash = hashlib.sha256(schema_sql.encode("utf-8")).hexdigest() - try: + with contextlib.suppress(Exception): # missing/corrupt cache → recompute below # Late import: resolves a test-patched hermes_constants.get_hermes_home. from hermes_constants import get_hermes_home as _home cache_path = _home() / "cache" / "schema_columns.json" @@ -539,8 +501,6 @@ class SessionSchemaMixin: isinstance(cols, dict) and all(isinstance(v, str) for v in cols.values()) for cols in tables.values() ): return tables - except Exception: - pass # missing/corrupt cache → recompute below ref = sqlite3.connect(":memory:") try: @@ -550,9 +510,8 @@ class SessionSchemaMixin: "SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%'" ).fetchall(): cols: Dict[str, str] = {} - for _cid, col_name, col_type, notnull, default, pk in ref.execute( - f'PRAGMA table_info("{tbl}")' - ).fetchall(): + info = ref.execute(f'PRAGMA table_info("{tbl}")').fetchall() + for _cid, col_name, col_type, notnull, default, pk in info: # Reconstruct the type expression for ALTER TABLE ADD COLUMN parts = [col_type] if col_type else [] if notnull and not pk: @@ -565,14 +524,12 @@ class SessionSchemaMixin: ref.close() if cache_path is not None: - try: + with contextlib.suppress(Exception): # cache write is best-effort cache_path.parent.mkdir(parents=True, exist_ok=True) fd, tmp = tempfile.mkstemp(dir=str(cache_path.parent), prefix=".schema_columns.") with os.fdopen(fd, "w", encoding="utf-8") as fh: json.dump({"schema_hash": schema_hash, "tables": table_columns}, fh) os.replace(tmp, cache_path) - except Exception: - pass # cache write is best-effort return table_columns def _reconcile_columns(self, cursor: sqlite3.Cursor) -> None: @@ -613,7 +570,7 @@ class SessionSchemaMixin: try: rows = cursor.execute(f'PRAGMA table_info("{table}")').fetchall() except sqlite3.OperationalError: - return None + rows = None if not rows: return None # row: (cid, name, type, notnull, dflt_value, pk) @@ -730,15 +687,12 @@ class SessionSchemaMixin: # Heal NULL ``active`` rows on every startup: older reconciler builds added ``active`` # without NOT NULL DEFAULT 1, so ``WHERE active = 1`` loaders hid whole histories. A # ``current_version < 12`` gate never re-ran for already-v12+ databases. - try: + with contextlib.suppress(sqlite3.OperationalError): cursor.execute("UPDATE messages SET active = 1 WHERE active IS NULL") - except sqlite3.OperationalError: - pass fts5_available = self._sqlite_supports_fts5(cursor) - self._fts_stale = cursor.execute( - "SELECT 1 FROM state_meta WHERE key = ? LIMIT 1", (FTS_STALE_KEY,) - ).fetchone() is not None + stale_row = cursor.execute("SELECT 1 FROM state_meta WHERE key = ? LIMIT 1", (FTS_STALE_KEY,)).fetchone() + self._fts_stale = stale_row is not None if self._fts_stale: # A prior process detached FTS after corruption; stay detached until a full rebuild. self._drop_all_fts_triggers(cursor) @@ -752,10 +706,9 @@ class SessionSchemaMixin: cursor.execute("INSERT INTO schema_version (version) VALUES (?)", (SCHEMA_VERSION,)) # Store provenance so fresh vs wiped stores are distinguishable. now_iso = datetime.datetime.now(datetime.timezone.utc).isoformat() - instance_id = str(uuid.uuid4()) cursor.executemany( "INSERT OR IGNORE INTO state_meta (key, value) VALUES (?, ?)", - [("store_instance_id", instance_id), ("store_created_at_utc", now_iso)], + [("store_instance_id", str(uuid.uuid4())), ("store_created_at_utc", now_iso)], ) else: self._run_data_migrations(cursor, row[0], fts5_available) @@ -773,7 +726,7 @@ class SessionSchemaMixin: # (v10 trigram backfill and v11 inline FTS re-index were superseded by v23 and removed.) if current_version < 16: # v16: tag delegate subagent rows so pickers stay clean after parent deletes orphan them. - try: + with contextlib.suppress(sqlite3.OperationalError): cursor.execute( "UPDATE sessions SET model_config = json_set(" "COALESCE(model_config, '{}'), '$._delegate_from', parent_session_id) " @@ -791,8 +744,6 @@ class SessionSchemaMixin: "AND NOT EXISTS (SELECT 1 FROM sessions ch " " WHERE ch.parent_session_id = sessions.id)" ) - except sqlite3.OperationalError: - pass if current_version < 18: # v18: best-effort gateway metadata backfill from sessions.json. try: @@ -801,10 +752,8 @@ class SessionSchemaMixin: logger.debug("v18 gateway metadata backfill skipped: %s", exc) if current_version < 20: # v20: seed session_model_usage from sessions aggregates (OR IGNORE: newer rows win). - try: + with contextlib.suppress(sqlite3.OperationalError): cursor.execute(_SESSION_MODEL_USAGE_V20_SEED_SQL) - except sqlite3.OperationalError: - pass if current_version < 22: self._migrate_v22_session_model_usage(cursor) # v23: FTS storage redesign (external-content tables). OPT-IN, NOT AUTOMATIC: the @@ -875,16 +824,14 @@ class SessionSchemaMixin: cursor.execute(_TITLE_UNIQUE_INDEX_SQL) except sqlite3.IntegrityError: try: - cursor.execute( - """UPDATE sessions AS older + cursor.execute("""UPDATE sessions AS older SET title = NULL WHERE title IS NOT NULL AND EXISTS ( SELECT 1 FROM sessions AS newer WHERE newer.title = older.title AND newer.rowid > older.rowid - )""" - ) + )""") logger.warning( "Cleared %d duplicate session title(s) while restoring the unique index", cursor.rowcount, ) @@ -905,12 +852,9 @@ class SessionSchemaMixin: # CJK was detached alongside the base indexes; its ensure path decides when it returns. self._ensure_fts_cjk_schema(cursor) else: - self._fts_enabled = False - self._trigram_available = False - self._fts_cjk_available = False + self._fts_enabled = self._trigram_available = self._fts_cjk_available = False else: base_sql, trigram_sql = _FTS_DDL[legacy_fts] - rebuild = self._rebuild_legacy_fts_indexes if legacy_fts else self._rebuild_fts_indexes # Measure BEFORE the DDL below runs (pre-repair state). Whether the trigram half is # creatable is only known AFTER _ensure_fts_schema, hence the halves combine at the `if`. base_triggers_missing = self._fts_triggers_missing(cursor, _FTS_BASE_TRIGGERS) @@ -922,7 +866,8 @@ class SessionSchemaMixin: self._trigram_available = trigram_enabled if base_triggers_missing or (trigram_enabled and trigram_triggers_missing): self._run_admitted_startup_rebuild( - cursor, lambda: rebuild(cursor, include_trigram=trigram_enabled), + cursor, + lambda: self._rebuild_fts_indexes(cursor, legacy=legacy_fts, include_trigram=trigram_enabled), ) if not legacy_fts: # CJK-bigram index: strictly additive, gated on the loadable tokenizer. @@ -950,9 +895,7 @@ class SessionSchemaMixin: cursor.execute(_STALE_KEY_UPSERT_SQL, (FTS_STALE_KEY,)) self._drop_all_fts_triggers(cursor) self._fts_stale = True - self._fts_enabled = False - self._trigram_available = False - self._fts_cjk_available = False + self._fts_enabled = self._trigram_available = self._fts_cjk_available = False def _backfill_gateway_metadata_from_sessions_json(self, cursor: sqlite3.Cursor) -> None: """One-time v18 backfill of gateway metadata from sessions.json. Only fills NULL @@ -986,13 +929,9 @@ class SessionSchemaMixin: END WHERE id = ?""", ( - entry.get("session_key") or key, - origin_dict.get("chat_id") if origin_dict is not None else None, - entry.get("chat_type"), - origin_dict.get("thread_id") if origin_dict is not None else None, - entry.get("display_name"), - json.dumps(origin) if origin_dict is not None else None, - 1 if entry.get("expiry_finalized") or entry.get("memory_flushed") else 0, - str(session_id), + entry.get("session_key") or key, origin_dict.get("chat_id") if origin_dict is not None else None, + entry.get("chat_type"), origin_dict.get("thread_id") if origin_dict is not None else None, + entry.get("display_name"), json.dumps(origin) if origin_dict is not None else None, + 1 if entry.get("expiry_finalized") or entry.get("memory_flushed") else 0, str(session_id), ), ) diff --git a/hermes_state_search.py b/hermes_state_search.py index 832602d58c..478420f0f9 100644 --- a/hermes_state_search.py +++ b/hermes_state_search.py @@ -4,6 +4,7 @@ Plain mixin for ``hermes_state.SessionDB`` (no ``__init__``/state of its own). Must never import hermes_state (cycle); shared constants live in hermes_state_common. """ +import contextlib import logging import re import sqlite3 @@ -143,8 +144,7 @@ def _search_select_sql(snippet_sql: str, from_sql: str, where: List[str], order_ def _search_filter_clauses( where: List[str], params: list, *, include_inactive: bool, source_filter: Optional[List[str]], - exclude_sources: Optional[List[str]], role_filter: Optional[List[str]], -) -> None: + exclude_sources: Optional[List[str]], role_filter: Optional[List[str]]) -> None: """Append the visibility/source/role predicates every search route shares. Live rows (active=1) AND compaction-archived rows (compacted=1) are discoverable; only rewind/undo rows (active=0, compacted=0) are hidden.""" @@ -204,9 +204,8 @@ class SessionSearchMixin: return self._rebuild_status("fts_cjk_rebuild") def _rebuild_status(self, prefix: str) -> Optional[Dict[str, Any]]: - rows = self._read_all( - "SELECT key, value FROM state_meta WHERE key IN (?, ?)", (f"{prefix}_high_water", f"{prefix}_progress"), - ) + rows = self._read_all("SELECT key, value FROM state_meta WHERE key IN (?, ?)", + (f"{prefix}_high_water", f"{prefix}_progress")) meta = {r["key"]: r["value"] for r in rows} high_water = meta.get(f"{prefix}_high_water") if high_water is None or int(high_water) <= 0: @@ -239,9 +238,8 @@ class SessionSearchMixin: def _fts_cjk_rebuild_finish(self) -> None: """Boundary sweep + clear the cjk markers; index becomes servable.""" - self._rebuild_finish("fts_cjk_rebuild", [ - self._BOUNDARY_SWEEP_SQL.format(table="messages_fts_cjk", extra="AND m.role <> 'tool' ") - ]) + sweep = self._BOUNDARY_SWEEP_SQL.format(table="messages_fts_cjk", extra="AND m.role <> 'tool' ") + self._rebuild_finish("fts_cjk_rebuild", [sweep]) self._fts_cjk_available = True logger.info("CJK FTS index backfill complete — serving CJK search.") @@ -265,19 +263,16 @@ class SessionSearchMixin: inserts = [self._CHUNK_INSERT_SQL.format(table="messages_fts", extra="")] if self._trigram_available: inserts.append(self._CHUNK_INSERT_SQL.format(table="messages_fts_trigram", extra=" AND role <> 'tool'")) - return self._rebuild_step( - "fts_rebuild", inserts, fail_msg="FTS rebuild chunk failed (will retry): %s", - finish=self._fts_rebuild_finish, - ) + return self._rebuild_step("fts_rebuild", inserts, fail_msg="FTS rebuild chunk failed (will retry): %s", + finish=self._fts_rebuild_finish) def fts_cjk_rebuild_step(self) -> bool: """Backfill one chunk of the CJK index. True while work remains.""" if not self._fts_enabled or not self._fts_cjk_loaded: return False - return self._rebuild_step( - "fts_cjk_rebuild", [self._CHUNK_INSERT_SQL.format(table="messages_fts_cjk", extra=" AND role <> 'tool'")], - fail_msg="CJK FTS rebuild chunk failed (will retry): %s", finish=self._fts_cjk_rebuild_finish, - ) + insert = self._CHUNK_INSERT_SQL.format(table="messages_fts_cjk", extra=" AND role <> 'tool'") + return self._rebuild_step("fts_cjk_rebuild", [insert], finish=self._fts_cjk_rebuild_finish, + fail_msg="CJK FTS rebuild chunk failed (will retry): %s") def _rebuild_step(self, prefix: str, insert_sqls: List[str], *, fail_msg: str, finish) -> bool: """Shared chunk engine for the base and CJK deferred backfills.""" @@ -322,13 +317,10 @@ class SessionSearchMixin: each chunk's scan is bounded (restarting the scan was O(n²)); compound-key tables keep the chunked ``LIMIT`` delete — they are small by construction.""" with self._lock: - trash = [ - r[0] for r in self._conn.execute( - "SELECT name FROM sqlite_master WHERE type = 'table' " - "AND name LIKE ? ESCAPE '\\'", - (self._FTS_TRASH_PREFIX.replace("_", "\\_") + "%",), - ).fetchall() - ] + trash = [r[0] for r in self._conn.execute( + "SELECT name FROM sqlite_master WHERE type = 'table' AND name LIKE ? ESCAPE '\\'", + (self._FTS_TRASH_PREFIX.replace("_", "\\_") + "%",), + ).fetchall()] if not trash: return False tbl = trash[0] @@ -358,9 +350,7 @@ class SessionSearchMixin: cur = conn.execute( f"DELETE FROM {tbl} WHERE ({key}) IN (SELECT {key} FROM {tbl} LIMIT {self._FTS_REBUILD_CHUNK_ROWS})" ) - if cur.rowcount == 0: - return _drop(conn) - return True # re-check: more trash tables / chunks may remain + return _drop(conn) if cur.rowcount == 0 else True # True: more trash tables / chunks may remain def _drop(conn, marker_key: Optional[str] = None) -> bool: """Drained — the DROP is cheap now. True: re-check for more trash.""" @@ -411,29 +401,22 @@ class SessionSearchMixin: except sqlite3.OperationalError: return False # table absent / FTS disabled mid-init — not this failure class - def _fts_index_known_empty(self, conn) -> bool: - """True when the base external-content index holds no rows; a missing table counts.""" - try: - return int(conn.execute("SELECT COUNT(*) FROM messages_fts_docsize").fetchone()[0]) == 0 - except sqlite3.OperationalError: - return True - - def _reset_fts_index_to_empty(self, conn) -> None: - """Truncate the v23 external-content tables via FTS5 ``'delete-all'`` (a plain DELETE is - O(rows) and corrupts the index when indexed rows diverged from ``messages``). The - backfill worker replays without an anti-join, so it needs a known-empty index.""" - for tbl in ("messages_fts", "messages_fts_trigram"): - try: - conn.execute(f"INSERT INTO {tbl}({tbl}) VALUES('delete-all')") - except sqlite3.OperationalError: - pass # table absent — already an empty surface - def _reseed_missing_progress(self, conn) -> None: """high_water without progress: fts_rebuild_step reads missing progress as "done by - another process" and optimize would no-op then stamp. Reset to known-empty, re-seed.""" + another process" and optimize would no-op then stamp. Reset to known-empty, re-seed. + Truncation goes through FTS5 ``'delete-all'`` (a plain DELETE is O(rows) and corrupts + the index when indexed rows diverged from ``messages``); the backfill worker replays + without an anti-join, so it needs a known-empty index. A missing docsize table counts + as empty.""" if _meta_row(conn, "fts_rebuild_progress") is None: - if not self._fts_index_known_empty(conn): - self._reset_fts_index_to_empty(conn) + try: + known_empty = int(conn.execute("SELECT COUNT(*) FROM messages_fts_docsize").fetchone()[0]) == 0 + except sqlite3.OperationalError: + known_empty = True + if not known_empty: + for tbl in ("messages_fts", "messages_fts_trigram"): + with contextlib.suppress(sqlite3.OperationalError): # table absent — already an empty surface + conn.execute(f"INSERT INTO {tbl}({tbl}) VALUES('delete-all')") self.set_meta("fts_rebuild_progress", "0", cursor=conn) def _seed_fts_rebuild_markers(self, conn, *, force: bool = False) -> int: @@ -494,26 +477,22 @@ class SessionSearchMixin: def _stage(conn): self._drop_fts_triggers(conn) conn.execute("DROP VIEW IF EXISTS messages_fts_trigram_src") - had = bool(conn.execute( + if conn.execute( "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name IN ('messages_fts', 'messages_fts_trigram') " "AND sql LIKE 'CREATE VIRTUAL TABLE%' LIMIT 1" - ).fetchone()) - if had: + ).fetchone(): conn.execute("PRAGMA writable_schema=ON") conn.execute( "DELETE FROM sqlite_master WHERE type = 'table' " "AND name IN ('messages_fts', 'messages_fts_trigram') AND sql LIKE 'CREATE VIRTUAL TABLE%'" ) conn.execute("PRAGMA writable_schema=RESET") - shadows = [ - r[0] for r in conn.execute( - "SELECT name FROM sqlite_master WHERE type = 'table' " - "AND (name LIKE 'messages_fts_%' ESCAPE '\\' " - "OR name LIKE 'messages_fts_trigram_%' ESCAPE '\\')" - ).fetchall() - ] - for sh in shadows: - conn.execute(f"ALTER TABLE {sh} RENAME TO fts_v22_trash_{sh}") + for row in conn.execute( + "SELECT name FROM sqlite_master WHERE type = 'table' " + "AND (name LIKE 'messages_fts_%' ESCAPE '\\' " + "OR name LIKE 'messages_fts_trigram_%' ESCAPE '\\')" + ).fetchall(): + conn.execute(f"ALTER TABLE {row[0]} RENAME TO fts_v22_trash_{row[0]}") # Claim the backfill BEFORE the empty v23 tables exist so a crash before # schema ensure resumes instead of stamping an empty index. hw = self._seed_fts_rebuild_markers(conn, force=True) @@ -536,17 +515,6 @@ class SessionSearchMixin: raise sqlite3.OperationalError(failure_message) self._conn.commit() - def _optimize_unsettled_reason(self, conn) -> Optional[str]: - """Refusal reason while optimize work remains, else None. An empty base index against - non-empty messages also refuses (settling there meant permanent search-index loss).""" - if _meta_row(conn, "fts_rebuild_high_water") is not None: - return "backfill_incomplete" - if self._has_fts_trash(conn): - return "teardown_incomplete" - if self._fts_external_index_empty_with_messages(conn): - return "backfill_incomplete" - return None - def _optimize_vacuum(self) -> bool: """Phase 3: reclaim freed pages to the OS. False when VACUUM failed (usually no free disk for its temp copy; a later VACUUM reclaims).""" @@ -570,10 +538,15 @@ class SessionSearchMixin: def _optimize_settle(self, conn) -> Optional[str]: """Phase 4 (inside the write transaction, so a concurrent writer cannot race a stamp past incomplete work): stamp the FTS layout (source of truth for "optimized"), clear the - "available" flag, advance a lagging schema_version. Returns a refusal reason or None.""" - refusal = self._optimize_unsettled_reason(conn) - if refusal is not None: - return refusal + "available" flag, advance a lagging schema_version. Returns a refusal reason or None. + Refuses while optimize work remains; an empty base index against non-empty messages + also refuses (settling there meant permanent search-index loss).""" + if _meta_row(conn, "fts_rebuild_high_water") is not None: + return "backfill_incomplete" + if self._has_fts_trash(conn): + return "teardown_incomplete" + if self._fts_external_index_empty_with_messages(conn): + return "backfill_incomplete" self.set_meta("fts_storage_version", str(FTS_STORAGE_VERSION), cursor=conn) _delete_meta(conn, "fts_optimize_available") conn.execute("UPDATE schema_version SET version = ? WHERE version < ?", (SCHEMA_VERSION, SCHEMA_VERSION)) @@ -612,10 +585,8 @@ class SessionSearchMixin: if progress_cb is None: return st = self.fts_rebuild_status() or self.fts_cjk_rebuild_status() - progress_cb({ - "phase": phase, "percent": st["percent"] if st else 100, - "indexed": st["indexed"] if st else 0, "total": st["total"] if st else 0, - }) + progress_cb({"phase": phase, "percent": st["percent"] if st else 100, + "indexed": st["indexed"] if st else 0, "total": st["total"] if st else 0}) def _drive(phase: str, step) -> None: """Run *step* to completion; the inter-chunk sleep is the single place the duty @@ -636,17 +607,14 @@ class SessionSearchMixin: # Phase 2: tear down the demoted legacy shadow tables in chunks. _emit("teardown") _drive("teardown", self._fts_teardown_trash_step) - with self._lock: still_pending = _meta_row(self._conn, "fts_rebuild_high_water") is not None still_trash = self._has_fts_trash(self._conn) empty_index = self._fts_external_index_empty_with_messages(self._conn) if still_pending or still_trash or empty_index: reason = "backfill_incomplete" if still_pending or empty_index else "teardown_incomplete" - logger.warning( - "FTS storage optimization did not settle (%s): pending=%s trash=%s empty_index=%s", - reason, still_pending, still_trash, empty_index, - ) + logger.warning("FTS storage optimization did not settle (%s): pending=%s trash=%s empty_index=%s", + reason, still_pending, still_trash, empty_index) return {"ok": False, "reason": reason, "vacuumed": None} vacuum_ok = None @@ -666,8 +634,7 @@ class SessionSearchMixin: def get_anchored_view( self, session_id: str, around_message_id: int, window: int = 5, bookend: int = 3, - keep_roles: Optional[Tuple[str, ...]] = ("user", "assistant"), - ) -> Dict[str, Any]: + keep_roles: Optional[Tuple[str, ...]] = ("user", "assistant")) -> Dict[str, Any]: """Anchored window (``get_messages_around``) plus session bookends, so one call yields the goal and the resolution of a long session. ``window`` is filtered to ``keep_roles`` (None disables) EXCEPT the anchor; ``bookend_start`` / ``bookend_end`` are the @@ -684,15 +651,11 @@ class SessionSearchMixin: if keep_roles is not None: keep_set = set(keep_roles) filtered_window = [m for m in window_rows if m.get("id") == around_message_id or m.get("role") in keep_set] - bookend_start_rows: List[Any] = [] bookend_end_rows: List[Any] = [] if bookend > 0: - role_clause = "" - role_params: list = [] - if keep_roles is not None: - role_clause = f" AND role IN ({','.join('?' for _ in keep_roles)})" - role_params = list(keep_roles) + role_clause = "" if keep_roles is None else f" AND role IN ({','.join('?' for _ in keep_roles)})" + role_params = [] if keep_roles is None else list(keep_roles) with self._read_ctx() as conn: def _bookend(op: str, boundary_id: int, order: str): return conn.execute( @@ -708,7 +671,6 @@ class SessionSearchMixin: def _hydrate(row) -> Dict[str, Any]: return self._row_to_message_dict(row, warn_context="get_anchored_view", summary_flag=False) - return { "window": filtered_window, "messages_before": primitive["messages_before"], "messages_after": primitive["messages_after"], @@ -717,8 +679,7 @@ class SessionSearchMixin: } def list_recent_user_messages( - self, session_id: str, limit: int = 20, include_inactive: bool = False - ) -> List[Dict[str, Any]]: + self, session_id: str, limit: int = 20, include_inactive: bool = False) -> List[Dict[str, Any]]: """The *limit* most-recent real user turns, newest first, as ``{id, timestamp, preview}`` (80 chars, whitespace collapsed); used by /rewind and ``/undo [N]``. Bookkeeping rows (``display_kind`` set) are excluded. Legacy compaction handoffs are role='user' rows @@ -727,17 +688,14 @@ class SessionSearchMixin: with a DB pick that includes them.""" active_clause = "" if include_inactive else " AND active = 1" display_clause = " AND (display_kind IS NULL OR display_kind = '')" - fetch_limit = int(limit) * 2 + 5 with self._lock: rows = self._conn.execute( "SELECT id, timestamp, content FROM messages WHERE session_id = ? AND role = 'user'" f"{active_clause}{display_clause} " "ORDER BY id DESC LIMIT ?", - (session_id, fetch_limit), + (session_id, int(limit) * 2 + 5), ).fetchall() - from agent.context_compressor import ContextCompressor - result: List[Dict[str, Any]] = [] for row in rows: if len(result) >= int(limit): @@ -745,8 +703,7 @@ class SessionSearchMixin: decoded = self._decode_content(row["content"]) if ContextCompressor._is_context_summary_content(decoded): continue # compaction handoff — never a user-originated turn - if isinstance(decoded, str): - # A /skill turn embeds the whole skill body; show what was typed. + if isinstance(decoded, str): # a /skill turn embeds the whole skill body; show what was typed preview = describe_skill_invocation(decoded) or decoded else: preview = _flatten_text(decoded) @@ -833,7 +790,8 @@ class SessionSearchMixin: def _trigram_route_ok(self, raw_query: str) -> bool: """Per-token CJK length gate for the trigram index: ``广西 OR 桂林 OR 漓江`` has 6 CJK chars total but 2 per token, so trigram returns 0.""" - return(self._count_cjk(raw_query) >= 3 and not self._has_short_cjk_token(raw_query) and self._trigram_available) + return (self._count_cjk(raw_query) >= 3 and not self._has_short_cjk_token(raw_query) + and self._trigram_available) def _describe_search_path(self, query: str) -> str: """Best-effort name of the routing path a query takes (log-only).""" @@ -848,26 +806,19 @@ class SessionSearchMixin: raw = sanitized.strip('"').strip() if self._fts_cjk_available and not self._has_lone_cjk_run(raw): return "fts_cjk" - if self._trigram_route_ok(raw): - return "trigram" - return "like_scan" + return "trigram" if self._trigram_route_ok(raw) else "like_scan" except Exception: return "unknown" # ── Query builders / runners ─────────────────────────────────────────── @staticmethod - def _fts_match_sql( - table: str, match_query: str, order_by_sql: str, *, include_inactive: bool, source_filter: Optional[List[str]], - exclude_sources: Optional[List[str]], role_filter: Optional[List[str]], limit: int, offset: int, - ) -> Tuple[str, list]: + def _fts_match_sql(table: str, match_query: str, order_by_sql: str, *, limit: int, offset: int, + **filters) -> Tuple[str, list]: """MATCH query + params against one FTS5 index joined to messages/sessions.""" where = [f"{table} MATCH ?"] params: list = [match_query] - _search_filter_clauses( - where, params, include_inactive=include_inactive, source_filter=source_filter, - exclude_sources=exclude_sources, role_filter=role_filter, - ) + _search_filter_clauses(where, params, **filters) params.extend([limit, offset]) sql = _search_select_sql( f"snippet({table}, -1, '>>>', '<<<', '...', 40) AS snippet", @@ -875,10 +826,8 @@ class SessionSearchMixin: ) return sql, params - def _match_rows( - self, table: str, match_query: str, order_by_sql: str, *, fail_open: Optional[str] = None, - operational_debug: Optional[str] = None, **kwargs, - ) -> Optional[List[Dict[str, Any]]]: + def _match_rows(self, table: str, match_query: str, order_by_sql: str, *, fail_open: Optional[str] = None, + operational_debug: Optional[str] = None, **kwargs) -> Optional[List[Dict[str, Any]]]: """Run one MATCH against *table*; ``None`` when the query cannot execute (tokenizer / syntax) so the caller falls back. *fail_open* names the index for the substring-capable routes: a corruption-class ``DatabaseError`` there detaches the @@ -896,8 +845,7 @@ class SessionSearchMixin: raise logger.warning( "%s FTS search hit a corruption error (%s); detached FTS and falling back to canonical LIKE.", - fail_open, exc, - ) + fail_open, exc) return None def _like_rows(self, where: List[str], params: list, *, order_by: str, limit_sql: str) -> List[Dict[str, Any]]: @@ -944,8 +892,7 @@ class SessionSearchMixin: return " OR ".join(compiled_groups), params, snippet_term def _search_messages_like_fallback( - self, query: str, *, limit: int, offset: int, sort: Optional[str], **filters - ) -> List[Dict[str, Any]]: + self, query: str, *, limit: int, offset: int, sort: Optional[str], **filters) -> List[Dict[str, Any]]: """Search canonical messages while derived FTS state is stale.""" predicate, params, snippet_term = self._compile_like_boolean_query(query) if not predicate or snippet_term is None: @@ -953,10 +900,8 @@ class SessionSearchMixin: where = [f"({predicate})"] _search_filter_clauses(where, params, **filters) order = "ASC" if isinstance(sort, str) and sort.strip().lower() == "oldest" else "DESC" - return self._like_rows( - where, [snippet_term, *params, limit, offset], - order_by=f"ORDER BY m.timestamp {order}, m.id {order}", limit_sql="LIMIT ? OFFSET ?", - ) + return self._like_rows(where, [snippet_term, *params, limit, offset], + order_by=f"ORDER BY m.timestamp {order}, m.id {order}", limit_sql="LIMIT ? OFFSET ?") def _refresh_fts_stale_state(self) -> None: """Observe fail-open initiated by another process sharing state.db.""" @@ -968,13 +913,10 @@ class SessionSearchMixin: return if stale is not None: self._fts_stale = True - self._fts_enabled = False - self._trigram_available = False - self._fts_cjk_available = False + self._fts_enabled = self._trigram_available = self._fts_cjk_available = False def _finalize_search_matches( - self, matches: List[Dict[str, Any]], result_fields: Optional[Collection[str]] = None - ) -> List[Dict[str, Any]]: + self, matches: List[Dict[str, Any]], result_fields: Optional[Collection[str]] = None) -> List[Dict[str, Any]]: """Attach neighboring messages (1 before + after, only when the projection consumes ``context``) and trim full content. Each context query takes its own read txn.""" if result_fields is None or "context" in result_fields: @@ -983,9 +925,8 @@ class SessionSearchMixin: with self._read_ctx() as conn: rows = conn.execute(_CONTEXT_WINDOW_SQL, (match["id"], match["id"])).fetchall() match["context"] = [ - {"role": row["role"], "content": _flatten_text(self._decode_content(row["content"]))[:200]} - for row in rows - ] + {"role": r["role"], "content": _flatten_text(self._decode_content(r["content"]))[:200]} + for r in rows] except Exception: match["context"] = [] # No route selects full content; the pop guards any future one that does. @@ -1009,16 +950,14 @@ class SessionSearchMixin: try: rows = self._search_messages_impl( query, source_filter=source_filter, exclude_sources=exclude_sources, role_filter=role_filter, - limit=limit, offset=offset, sort=sort, include_inactive=include_inactive, fields=fields, - ) + limit=limit, offset=offset, sort=sort, include_inactive=include_inactive, fields=fields) return rows finally: elapsed_ms = (time.time() - started) * 1000.0 if elapsed_ms >= env_float("HERMES_SEARCH_SLOW_MS", 1000.0): - logger.info( - "slow session search: path=%s elapsed=%.0fms rows=%s query=%r", self._describe_search_path(query), - elapsed_ms, len(rows) if rows is not None else "err", query[: 200], - ) + logger.info("slow session search: path=%s elapsed=%.0fms rows=%s query=%r", + self._describe_search_path(query), elapsed_ms, len(rows) if rows is not None else "err", + query[: 200]) def _search_messages_impl( self, query: str, source_filter: List[str] = None, exclude_sources: List[str] = None, @@ -1036,11 +975,8 @@ class SessionSearchMixin: query = self._sanitize_fts5_query(query) if not query: return [] - - filters = dict( - include_inactive=include_inactive, source_filter=source_filter, - exclude_sources=exclude_sources, role_filter=role_filter, - ) + filters = dict(include_inactive=include_inactive, source_filter=source_filter, + exclude_sources=exclude_sources, role_filter=role_filter) self._refresh_fts_stale_state() if self._fts_stale: matches = self._search_messages_like_fallback(query, limit=limit, offset=offset, sort=sort, **filters) @@ -1103,8 +1039,7 @@ class SessionSearchMixin: if self._fts_cjk_available and not wants_tool_rows and not self._has_lone_cjk_run(raw_query): matches = self._match_rows( "messages_fts_cjk", match_query, fail_open="CJK-bigram", - operational_debug="messages_fts_cjk query failed; falling back to trigram/LIKE", **route, - ) + operational_debug="messages_fts_cjk query failed; falling back to trigram/LIKE", **route) if matches is not None: return matches if self._trigram_route_ok(raw_query) and not wants_tool_rows: @@ -1112,17 +1047,13 @@ class SessionSearchMixin: if matches is not None: return matches non_op_tokens = _non_operator_tokens(raw_query) or [raw_query] - like_params: list = [] - for tok in non_op_tokens: - like_params += _like_params(tok) + like_params: list = [p for tok in non_op_tokens for p in _like_params(tok)] like_where = [f"({' OR '.join([_LIKE_ANY_COLUMN_SQL] * len(non_op_tokens))})"] filters = {k: route[k] for k in ("include_inactive", "source_filter", "exclude_sources", "role_filter")} _search_filter_clauses(like_where, like_params, **filters) # instr() for the snippet uses the first search token. - return self._like_rows( - like_where, [non_op_tokens[0], *like_params, route["limit"], route["offset"]], - order_by="ORDER BY m.timestamp DESC", limit_sql="LIMIT ? OFFSET ?", - ) + return self._like_rows(like_where, [non_op_tokens[0], *like_params, route["limit"], route["offset"]], + order_by="ORDER BY m.timestamp DESC", limit_sql="LIMIT ? OFFSET ?") def _search_unindexed_gap(self, fts_query: str, limit: int, **filters) -> List[Dict[str, Any]]: """LIKE-scan ids in (fts_rebuild_progress, fts_rebuild_high_water] — rows the deferred @@ -1131,26 +1062,19 @@ class SessionSearchMixin: status = self.fts_rebuild_status() if status is None or limit <= 0: return [] - terms = [ - tok for tok in (t.strip('"').strip("*").strip() for t in _LIKE_TOKEN_RE.findall(fts_query)) - if tok and tok.upper() not in _LIKE_SKIP_TOKENS - ] + terms = [tok for tok in (t.strip('"').strip("*").strip() for t in _LIKE_TOKEN_RE.findall(fts_query)) + if tok and tok.upper() not in _LIKE_SKIP_TOKENS] if not terms: return [] - where = ["m.id > ? AND m.id <= ?"] - params: list = [status["indexed"], status["total"]] - for term in terms: - where.append(_LIKE_ANY_COLUMN_SQL) - params += _like_params(term) + where = ["m.id > ? AND m.id <= ?", *([_LIKE_ANY_COLUMN_SQL] * len(terms))] + params: list = [status["indexed"], status["total"], *(p for term in terms for p in _like_params(term))] _search_filter_clauses(where, params, **filters) - return self._like_rows( - where, [terms[0], *params, limit], order_by="ORDER BY m.timestamp DESC", limit_sql="LIMIT ?", - ) + return self._like_rows(where, [terms[0], *params, limit], order_by="ORDER BY m.timestamp DESC", + limit_sql="LIMIT ?") def search_sessions_by_id( self, query: str, limit: int = 20, include_archived: bool = True, source: str = None, - sources: List[str] = None, exclude_sources: List[str] = None, - ) -> List[Dict[str, Any]]: + sources: List[str] = None, exclude_sources: List[str] = None) -> List[Dict[str, Any]]: """Search surfaced sessions by exact/prefix/substring session id. Also matches ``_lineage_root_id`` so an old compression root id resolves to the live continuation.""" needle = (query or "").strip().lower() @@ -1160,18 +1084,13 @@ class SessionSearchMixin: # chain) into SQL; over-fetch so the in-Python ranking has candidates. candidates = self.list_sessions_rich( source=source, sources=sources, exclude_sources=exclude_sources, limit=max(limit * 4, limit), - offset=0, include_archived=include_archived, order_by_last_active=True, id_query=needle, - ) + offset=0, include_archived=include_archived, order_by_last_active=True, id_query=needle) def score(row: Dict[str, Any]) -> int: - ids = [str(row.get("id") or ""), str(row.get("_lineage_root_id") or "")] - normalized = [value.lower() for value in ids if value] + normalized = [v.lower() for v in (str(row.get("id") or ""), str(row.get("_lineage_root_id") or "")) if v] if any(value == needle for value in normalized): return 0 - if any(value.startswith(needle) for value in normalized): - return 1 - return 2 - + return 1 if any(value.startswith(needle) for value in normalized) else 2 ranked = sorted(enumerate(candidates), key=lambda item: (score(item[1]), item[0])) return [row for _, row in ranked[:limit]] @@ -1214,8 +1133,7 @@ class SessionSearchMixin: with fts_rebuild_admission(self.db_path) as admitted: if not admitted: logger.warning( - "Deferred in-place FTS rebuild: another process holds the rebuild authority for this state.db." - ) + "Deferred in-place FTS rebuild: another process holds the rebuild authority for this state.db.") return 0 with self._lock: for tbl in self._present_fts_tables(): @@ -1243,7 +1161,6 @@ class SessionSearchMixin: if max_commands is None: max_commands = self._FTS_MERGE_COMMANDS_PER_PASS _positive_int("max_commands", max_commands) - executed = 0 with self._lock: for tbl in self._present_fts_tables(): diff --git a/hermes_state_telegram.py b/hermes_state_telegram.py index a7172a3370..c8a37d97ae 100644 --- a/hermes_state_telegram.py +++ b/hermes_state_telegram.py @@ -2,6 +2,7 @@ from __future__ import annotations +import contextlib import logging import sqlite3 import time @@ -99,20 +100,12 @@ class SessionTelegramTopicsMixin: tables (nobody ran ``/topic``) by returning their empty value; only ``enable``/``bind`` run the migration.""" - def _topic_read_one(self, sql: str, params, default=None): - """``fetchone`` that treats an unmigrated table as *default*.""" - return self._topic_read(self._read_one, sql, params, default) - - def _topic_read_all(self, sql: str, params) -> list: - """``fetchall`` that treats an unmigrated table as no rows.""" - return self._topic_read(self._read_all, sql, params, []) - - @staticmethod - def _topic_read(reader, sql: str, params, default): + def _topic_read_one(self, sql: str, params): + """``fetchone`` that treats an unmigrated table as None.""" try: - return reader(sql, params) + return self._read_one(sql, params) except sqlite3.OperationalError: - return default + return None def apply_telegram_topic_migration(self) -> None: """Create Telegram DM topic-mode tables on explicit /topic opt-in. Deliberately NOT @@ -129,25 +122,21 @@ class SessionTelegramTopicsMixin: # v1/v2 → v3. SQLite can't ALTER a PK or FK, so rebuild (also supplies v2's # ON DELETE CASCADE). Legacy rows land in "default" only. legacy_columns = columns.replace("profile_name, ", "", 1) - conn.executescript( - f""" + conn.executescript(f""" CREATE TABLE {table}_new ({ddl}); INSERT INTO {table}_new ({columns}) SELECT 'default', {legacy_columns} FROM {table}; DROP TABLE {table}; ALTER TABLE {table}_new RENAME TO {table}; - """ - ) + """) # Indexes after any rebuild: the user index needs profile_name. - conn.executescript( - """ + conn.executescript(""" CREATE UNIQUE INDEX IF NOT EXISTS idx_telegram_dm_topic_bindings_session ON telegram_dm_topic_bindings(session_id); CREATE INDEX IF NOT EXISTS idx_telegram_dm_topic_bindings_user ON telegram_dm_topic_bindings(profile_name, user_id, chat_id); - """ - ) + """) conn.execute( "INSERT INTO state_meta (key, value) VALUES (?, ?) " "ON CONFLICT(key) DO UPDATE SET value = excluded.value", @@ -168,8 +157,7 @@ class SessionTelegramTopicsMixin: def _to_int(value: Optional[bool]) -> Optional[int]: return None if value is None else (1 if value else 0) - self._write_sql( - """ + self._write_sql(""" INSERT INTO telegram_dm_topic_mode ( profile_name, chat_id, user_id, enabled, activated_at, updated_at, has_topics_enabled, allows_users_to_create_topics, @@ -182,10 +170,8 @@ class SessionTelegramTopicsMixin: has_topics_enabled = excluded.has_topics_enabled, allows_users_to_create_topics = excluded.allows_users_to_create_topics, capability_checked_at = excluded.capability_checked_at - """, - (profile_name, str(chat_id), str(user_id), now, now, - _to_int(has_topics_enabled), _to_int(allows_users_to_create_topics), now), - ) + """, (profile_name, str(chat_id), str(user_id), now, now, + _to_int(has_topics_enabled), _to_int(allows_users_to_create_topics), now)) def disable_telegram_topic_mode( self, *, chat_id: str, profile_name: str = "default", clear_bindings: bool = True @@ -196,7 +182,7 @@ class SessionTelegramTopicsMixin: profile_name = _normalize_telegram_topic_profile_name(profile_name) def _do(conn): - try: + with contextlib.suppress(sqlite3.OperationalError): conn.execute( "UPDATE telegram_dm_topic_mode SET enabled = 0, updated_at = ? " "WHERE profile_name = ? AND chat_id = ?", @@ -207,20 +193,15 @@ class SessionTelegramTopicsMixin: "DELETE FROM telegram_dm_topic_bindings WHERE profile_name = ? AND chat_id = ?", (profile_name, str(chat_id)), ) - except sqlite3.OperationalError: - return self._execute_write(_do) def is_telegram_topic_mode_enabled(self, *, chat_id: str, user_id: str, profile_name: str = "default") -> bool: """Return whether Telegram DM topic mode is enabled for this chat/user.""" profile_name = _normalize_telegram_topic_profile_name(profile_name) - row = self._topic_read_one( - """ + row = self._topic_read_one(""" SELECT enabled FROM telegram_dm_topic_mode WHERE profile_name = ? AND chat_id = ? AND user_id = ? - """, - (profile_name, str(chat_id), str(user_id)), - ) + """, (profile_name, str(chat_id), str(user_id))) return bool(row[0]) if row is not None else False def get_telegram_topic_binding( @@ -228,13 +209,10 @@ class SessionTelegramTopicsMixin: ) -> Optional[Dict[str, Any]]: """Return the session binding for a Telegram DM topic, if present.""" profile_name = _normalize_telegram_topic_profile_name(profile_name) - row = self._topic_read_one( - """ + row = self._topic_read_one(""" SELECT * FROM telegram_dm_topic_bindings WHERE profile_name = ? AND chat_id = ? AND thread_id = ? - """, - (profile_name, str(chat_id), str(thread_id)), - ) + """, (profile_name, str(chat_id), str(thread_id))) return dict(row) if row else None def list_telegram_topic_bindings_for_chat( @@ -242,21 +220,21 @@ class SessionTelegramTopicsMixin: ) -> List[Dict[str, Any]]: """All bindings for one chat, newest first ([] when the table is absent).""" profile_name = _normalize_telegram_topic_profile_name(profile_name) - rows = self._topic_read_all( - "SELECT * FROM telegram_dm_topic_bindings WHERE profile_name = ? AND chat_id = ? ORDER BY updated_at DESC", - (profile_name, str(chat_id)), - ) + try: + rows = self._read_all( + "SELECT * FROM telegram_dm_topic_bindings WHERE profile_name = ? AND chat_id = ? ORDER BY updated_at DESC", + (profile_name, str(chat_id)), + ) + except sqlite3.OperationalError: + return [] return [dict(row) for row in rows] def get_telegram_topic_binding_by_session(self, *, session_id: str) -> Optional[Dict[str, Any]]: """Reverse lookup via the UNIQUE INDEX on session_id; None when unbound.""" - row = self._topic_read_one( - """ + row = self._topic_read_one(""" SELECT * FROM telegram_dm_topic_bindings WHERE session_id = ? - """, - (str(session_id),), - ) + """, (str(session_id),)) return dict(row) if row else None def delete_telegram_topic_binding(self, *, chat_id: str, thread_id: str, profile_name: str = "default") -> int: @@ -268,41 +246,32 @@ class SessionTelegramTopicsMixin: transaction, or a user who disabled topics in the Telegram client (not via ``/topic off``) stays stuck. Returns the number of rows deleted; absent binding or unmigrated tables are silent no-ops (never raise from a cleanup hot path).""" - chat_id = str(chat_id) - thread_id = str(thread_id) + chat_id, thread_id = str(chat_id), str(thread_id) profile_name = _normalize_telegram_topic_profile_name(profile_name) def _do(conn) -> int: try: - deleted = conn.execute( - """ + deleted = conn.execute(""" DELETE FROM telegram_dm_topic_bindings WHERE profile_name = ? AND chat_id = ? AND thread_id = ? - """, - (profile_name, chat_id, thread_id), - ).rowcount or 0 + """, (profile_name, chat_id, thread_id)).rowcount or 0 except sqlite3.OperationalError: return 0 if not deleted: return 0 # Last binding gone → disable topic mode in the same transaction (no - # read-after-prune race). - try: - remaining = conn.execute( - """ + # read-after-prune race). telegram_dm_topic_mode absent — binding prune still stands. + with contextlib.suppress(sqlite3.OperationalError): + remaining = conn.execute(""" SELECT 1 FROM telegram_dm_topic_bindings WHERE profile_name = ? AND chat_id = ? LIMIT 1 - """, - (profile_name, chat_id), - ).fetchone() + """, (profile_name, chat_id)).fetchone() if remaining is None: conn.execute( "UPDATE telegram_dm_topic_mode SET enabled = 0, updated_at = ? " "WHERE profile_name = ? AND chat_id = ?", (time.time(), profile_name, chat_id), ) - except sqlite3.OperationalError: - pass # telegram_dm_topic_mode absent — binding prune still stands. return deleted return self._execute_write(_do) @@ -321,24 +290,16 @@ class SessionTelegramTopicsMixin: profile_name = _normalize_telegram_topic_profile_name(profile_name) def _do(conn): - existing_session = conn.execute( - """ + existing_session = conn.execute(""" SELECT profile_name, chat_id, thread_id FROM telegram_dm_topic_bindings WHERE session_id = ? - """, - (session_id,), - ).fetchone() + """, (session_id,)).fetchone() if existing_session is not None: linked_profile, linked_chat, linked_thread = existing_session - if ( - str(linked_profile) != profile_name - or str(linked_chat) != chat_id - or str(linked_thread) != thread_id - ): + if (str(linked_profile), str(linked_chat), str(linked_thread)) != (profile_name, chat_id, thread_id): raise ValueError("session is already linked to another Telegram topic") - conn.execute( - """ + conn.execute(""" INSERT INTO telegram_dm_topic_bindings ( profile_name, chat_id, thread_id, user_id, session_key, session_id, managed_mode, linked_at, updated_at @@ -349,22 +310,16 @@ class SessionTelegramTopicsMixin: session_id = excluded.session_id, managed_mode = excluded.managed_mode, updated_at = excluded.updated_at - """, - (profile_name, chat_id, thread_id, user_id, session_key, session_id, - managed_mode, now, now), - ) + """, (profile_name, chat_id, thread_id, user_id, session_key, session_id, managed_mode, now, now)) self._execute_write(_do) def is_telegram_session_linked_to_topic(self, *, session_id: str) -> bool: """True if the session is bound to any Telegram DM topic (absent tables → False).""" - row = self._topic_read_one( - """ + row = self._topic_read_one(""" SELECT 1 FROM telegram_dm_topic_bindings WHERE session_id = ? LIMIT 1 - """, - (str(session_id),), - ) + """, (str(session_id),)) return row is not None def list_unlinked_telegram_sessions_for_user( diff --git a/hermes_state_titles.py b/hermes_state_titles.py index 0297ddf11b..5925593dba 100644 --- a/hermes_state_titles.py +++ b/hermes_state_titles.py @@ -27,9 +27,8 @@ class SessionTitlesMixin: def _title_rank(cls, source: Optional[str]) -> int: """Rank a stored title_source. NULL (pre-provenance rows) is indistinguishable from a manual ``/title`` of that era, so it ranks as ``user``.""" - if source is None: - return cls._TITLE_SOURCE_RANK[cls.TITLE_SOURCE_USER] - return cls._TITLE_SOURCE_RANK.get(str(source), 0) + rank = cls._TITLE_SOURCE_RANK + return rank[cls.TITLE_SOURCE_USER] if source is None else rank.get(str(source), 0) @staticmethod def sanitize_title(title: Optional[str]) -> Optional[str]: @@ -53,8 +52,7 @@ class SessionTitlesMixin: if not ancestor_id or not descendant_id or ancestor_id == descendant_id: return False edge = _COMPRESSION_CHILD_SQL.format(a="child") - row = conn.execute( - f""" + return conn.execute(f""" WITH RECURSIVE ancestors(id) AS ( SELECT ? UNION @@ -65,10 +63,7 @@ class SessionTitlesMixin: WHERE {edge} ) SELECT 1 FROM ancestors WHERE id = ? AND id != ? LIMIT 1 - """, - (descendant_id, ancestor_id, descendant_id), - ).fetchone() - return row is not None + """, (descendant_id, ancestor_id, descendant_id)).fetchone() is not None def _set_session_title(self, session_id: str, title: str, *, source: str) -> bool: """Write a title, enforcing provenance precedence. A ``user`` write always lands; @@ -91,17 +86,12 @@ class SessionTitlesMixin: # exact-title lookup on every open), so a rename orphans the conversation. Hidden # is the discriminator: canonical chats are born hidden; a visible session merely # named "Bot Chat" stays renameable. Provenance-blind. - if ( - (current["title"] or "") == self.CANONICAL_BOT_CHAT_TITLE - and bool(current["hidden"]) - and title != self.CANONICAL_BOT_CHAT_TITLE - ): + if ((current["title"] or "") == self.CANONICAL_BOT_CHAT_TITLE and bool(current["hidden"]) + and title != self.CANONICAL_BOT_CHAT_TITLE): if is_user: - raise ValueError( - "This is the bot's canonical Bot Chat — its name is its " - "identity, and renaming it would orphan the conversation. " - "To start fresh, create a new bot instead." - ) + raise ValueError("This is the bot's canonical Bot Chat — its name is its " + "identity, and renaming it would orphan the conversation. " + "To start fresh, create a new bot instead.") return 0 if not is_user and current["title"] is not None and self._title_rank(current["title_source"]) >= new_rank: return 0 @@ -119,11 +109,10 @@ class SessionTitlesMixin: raise ValueError(f"Title '{title}' is already in use by session {conflict_id}") # CAS on the values just read (``IS`` is NULL-safe): a concurrent write between # the SELECT and here loses instead of being overwritten. - cursor = conn.execute( + return conn.execute( "UPDATE sessions SET title = ?, title_source = ? WHERE id = ? AND title IS ? AND title_source IS ?", (title, source if title else None, session_id, current["title"], current["title_source"]), - ) - return cursor.rowcount + ).rowcount return self._execute_write(_do) > 0 @@ -150,9 +139,7 @@ class SessionTitlesMixin: def get_session_title_source(self, session_id: str) -> Optional[str]: """Get the provenance of a session's title, or None when untitled.""" row = self._read_one("SELECT title, title_source FROM sessions WHERE id = ?", (session_id,)) - if not row or row["title"] is None: - return None - return row["title_source"] + return row["title_source"] if row and row["title"] is not None else None def set_session_title_source(self, session_id: str, source: str) -> bool: """Overwrite a title's provenance without touching the text (a title copied across a @@ -179,9 +166,7 @@ class SessionTitlesMixin: "SELECT id, title, started_at FROM sessions " "WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC", (f"{_escape_like(title)} #%",)) - if numbered: - return numbered[0]["id"] - return exact["id"] if exact else None + return numbered[0]["id"] if numbered else (exact["id"] if exact else None) def get_next_title_in_lineage(self, base_title: str) -> str: """Next title in a lineage ("my session" -> "my session #2"): strip any " #N" suffix, diff --git a/hermes_state_usage.py b/hermes_state_usage.py index 61396dc6b8..e9ff699112 100644 --- a/hermes_state_usage.py +++ b/hermes_state_usage.py @@ -4,6 +4,7 @@ per-model usage rows, and billing-route columns. Writer thread state lives on th from __future__ import annotations import atexit +import contextlib import logging import threading import time @@ -91,16 +92,13 @@ class SessionUsageMixin: self.flush_token_counts() def _do(conn): - conn.execute( - """UPDATE sessions SET + conn.execute("""UPDATE sessions SET billing_provider = ?, billing_base_url = ?, billing_mode = COALESCE(?, billing_mode), system_prompt = NULL, system_prompt_hash = NULL - WHERE id = ?""", - (provider, base_url, billing_mode, session_id), - ) + WHERE id = ?""", (provider, base_url, billing_mode, session_id)) self._delete_unreferenced_system_prompts(conn) self._execute_write(_do) @@ -217,7 +215,7 @@ class SessionUsageMixin: for session_id, kwargs in batch: key = None if not kwargs.get("absolute"): - key = (session_id,) + tuple(kwargs.get(f) for f in self._TOKEN_DELTA_ROUTE_FIELDS) + key = (session_id, *(kwargs.get(f) for f in self._TOKEN_DELTA_ROUTE_FIELDS)) if groups and key is not None and groups[-1][0] == key: merged = groups[-1][2] for f in self._TOKEN_DELTA_SUM_FIELDS: @@ -268,10 +266,8 @@ class SessionUsageMixin: self._apply_claimed_batch(batch) def _drain_token_queue_at_exit(self) -> None: - try: + with contextlib.suppress(Exception): # never fatal at interpreter shutdown self._stop_token_writer() - except Exception: - pass # never fatal at interpreter shutdown def update_token_counts( self, session_id: str, input_tokens: int=0, output_tokens: int=0, model: str=None, cache_read_tokens: int=0, @@ -288,10 +284,8 @@ class SessionUsageMixin: # locking, and the UPDATE would silently affect 0 rows. self._insert_session_row(session_id, "unknown", model=model) sql = _TOKEN_UPDATE_ABSOLUTE_SQL if absolute else _TOKEN_UPDATE_DELTA_SQL - has_usage = bool( - input_tokens or output_tokens or cache_read_tokens - or cache_write_tokens or reasoning_tokens or api_call_count or estimated_cost_usd - ) + has_usage = bool(input_tokens or output_tokens or cache_read_tokens or cache_write_tokens or reasoning_tokens + or api_call_count or estimated_cost_usd) has_accounted_usage = bool(has_usage or actual_cost_usd) params = ( input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens, @@ -320,13 +314,10 @@ class SessionUsageMixin: and (existing.get("model") != model or existing.get("billing_provider") != billing_provider) ) if first_accounted_route: - conn.execute( - """UPDATE sessions + conn.execute("""UPDATE sessions SET model = ?, billing_provider = ?, billing_base_url = ?, billing_mode = ? - WHERE id = ?""", - (model, billing_provider, billing_base_url, billing_mode, session_id), - ) + WHERE id = ?""", (model, billing_provider, billing_base_url, billing_mode, session_id)) conn.execute(sql, params) if record_model_usage: self._record_model_usage(conn, session_id, **usage) @@ -350,16 +341,12 @@ class SessionUsageMixin: sess = dict(row) if (row is not None and not task) else {} counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)] now = time.time() - conn.execute( - _MODEL_USAGE_UPSERT_SQL, - ( - session_id, model or sess.get("model") or "unknown", - billing_provider or sess.get("billing_provider") or "", - billing_base_url or sess.get("billing_base_url") or "", - billing_mode or sess.get("billing_mode") or "", task or "", - api_call_count or 0, *counts, - float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0), - cost_status, cost_source, now, now)) + conn.execute(_MODEL_USAGE_UPSERT_SQL, ( + session_id, model or sess.get("model") or "unknown", + billing_provider or sess.get("billing_provider") or "", + billing_base_url or sess.get("billing_base_url") or "", + billing_mode or sess.get("billing_mode") or "", task or "", api_call_count or 0, *counts, + float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0), cost_status, cost_source, now, now)) def record_auxiliary_usage( self, session_id: str, task: str, *, model: Optional[str]=None, billing_provider: Optional[str]=None, @@ -386,13 +373,10 @@ class SessionUsageMixin: params: List[Any] = [min_message_count] if not include_archived: where.append("COALESCE(archived, 0) = 0") - row = self._read_one( - f""" + row = self._read_one(f""" SELECT COALESCE(SUM(COALESCE(input_tokens, 0) + COALESCE(output_tokens, 0)), 0), COALESCE(SUM(COALESCE(actual_cost_usd, estimated_cost_usd, 0)), 0) FROM sessions WHERE {' AND '.join(where)} - """, - params, - ) + """, params) return {"tokens": int(row[0] or 0), "cost_usd": float(row[1] or 0.0)}