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:
+25
-54
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)}
|
||||||
|
|||||||
Reference in New Issue
Block a user