refactor(state): search/schema/registry/portability/telegram/usage/titles — inline single-use helpers, contextlib.suppress ladders, pack wrappers around unchanged SQL literals

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