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