refactor(state): table-drive parent-metadata inheritance SQL, unify end-stamp bump and delete cascades, collapse defensive layers in hermes_state_sessions
This commit is contained in:
@@ -1108,7 +1108,8 @@ class SessionDB(
|
||||
<3.12 has no setconfig; the residual checkpoint only carries
|
||||
pre-quarantine committed frames, which is tolerable."""
|
||||
flag = getattr(sqlite3, "SQLITE_DBCONFIG_NO_CKPT_ON_CLOSE", None)
|
||||
setconfig = getattr(self._conn, "setconfig", None)
|
||||
conn = self._conn
|
||||
setconfig = getattr(conn, "setconfig", None)
|
||||
if flag is None or setconfig is None:
|
||||
return
|
||||
try:
|
||||
|
||||
@@ -18,6 +18,7 @@ from hermes_state_common import (
|
||||
_LISTABLE_CHILD_SQL, _PREVIEW_ELIGIBLE_SQL, _PREVIEW_RAW_SELECT, _RECOVERABLE_END_REASONS,
|
||||
_RECOVERABLE_END_REASONS_SQL, _RESET_END_REASONS, _legacy_reset_child_sql, _shape_preview,
|
||||
_sql_session_last_active, _sql_session_last_active_by_id, escape_like as _escape_like,
|
||||
_placeholders as _session_ids_placeholders,
|
||||
)
|
||||
|
||||
# caplog tests pin the "hermes_state" logger name.
|
||||
@@ -43,13 +44,10 @@ def _parse_model_config(raw: Any) -> Dict[str, Any]:
|
||||
"""Tolerant ``model_config`` decode: JSON text or dict -> dict copy; anything else -> {}."""
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
try:
|
||||
parsed = json.loads(raw)
|
||||
raw = json.loads(raw)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
return {}
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
if isinstance(raw, dict):
|
||||
return dict(raw)
|
||||
return {}
|
||||
return dict(raw) if isinstance(raw, dict) else {}
|
||||
|
||||
|
||||
def _cwd_prefix_clause(cwd_prefix: str) -> Tuple[str, List[str]]:
|
||||
@@ -86,10 +84,6 @@ _PREVIEW_COL_SQL = f"""COALESCE(
|
||||
) AS _preview_raw"""
|
||||
|
||||
|
||||
def _session_ids_placeholders(ids) -> str:
|
||||
return ",".join("?" * len(ids))
|
||||
|
||||
|
||||
def _where_sql(clauses: List[str], lead: str = "") -> str:
|
||||
"""``WHERE a AND b`` (with *lead* prefix) or "" when there are no clauses."""
|
||||
return f"{lead}WHERE {' AND '.join(clauses)}" if clauses else ""
|
||||
@@ -161,14 +155,11 @@ def _delete_delegate_children(conn, parent_ids: List[str]) -> List[str]:
|
||||
ph = _session_ids_placeholders(ids)
|
||||
conn.execute(f"DELETE FROM messages WHERE session_id IN ({ph})", ids)
|
||||
# FK safety: orphan any untagged stragglers pointing at a doomed row.
|
||||
conn.execute(
|
||||
f"UPDATE sessions SET parent_session_id = NULL WHERE parent_session_id IN ({ph})", ids,
|
||||
)
|
||||
conn.execute(f"UPDATE sessions SET parent_session_id = NULL WHERE parent_session_id IN ({ph})", ids)
|
||||
conn.execute(f"DELETE FROM sessions WHERE id IN ({ph})", ids)
|
||||
return ids
|
||||
|
||||
|
||||
|
||||
# Lifecycle statuses surfaced by session pickers; classified from the final
|
||||
# message row ONLY so it stays O(1) per session.
|
||||
SESSION_STATUS_COMPLETE = "complete"
|
||||
@@ -189,9 +180,7 @@ def classify_session_status(
|
||||
if (finish_reason or "").strip().lower() in _ERROR_FINISH_REASONS:
|
||||
return SESSION_STATUS_ERROR
|
||||
r = (role or "").strip().lower()
|
||||
if r == "assistant":
|
||||
return SESSION_STATUS_INTERRUPTED if has_tool_calls else SESSION_STATUS_COMPLETE
|
||||
if r in {"user", "tool"}:
|
||||
if r in {"user", "tool"} or (r == "assistant" and has_tool_calls):
|
||||
return SESSION_STATUS_INTERRUPTED
|
||||
return SESSION_STATUS_COMPLETE
|
||||
|
||||
@@ -205,6 +194,38 @@ _SAME_KEY_NAMESPACE_SQL = (
|
||||
)
|
||||
|
||||
|
||||
def _inherit_col_sql(col: str, extra: str = "") -> str:
|
||||
"""``col = COALESCE(sessions.col, (SELECT p.col FROM parent))`` (whitespace is part of the SQL text)."""
|
||||
pad = " " * (30 + len(col))
|
||||
return (
|
||||
f"{col} = COALESCE(sessions.{col},\n{pad}(SELECT p.{col} FROM sessions p\n"
|
||||
f"{pad} WHERE p.id = sessions.parent_session_id{extra}))"
|
||||
)
|
||||
|
||||
|
||||
_INHERIT_SEP = ",\n" + " " * 27
|
||||
_INHERIT_PARENT_META_SQL = (
|
||||
"UPDATE sessions\n SET "
|
||||
+ _INHERIT_SEP.join((
|
||||
*(_inherit_col_sql(c) for c in ("cwd", "git_repo_root", "git_branch")),
|
||||
_inherit_col_sql("profile_name", "\n" + " " * 46 + f"AND ({_SAME_KEY_NAMESPACE_SQL})"),
|
||||
))
|
||||
+ "\n WHERE id = ? AND parent_session_id IS NOT NULL"
|
||||
)
|
||||
_INHERIT_PARENT_ROUTING_SQL = (
|
||||
"UPDATE sessions\n SET "
|
||||
+ _INHERIT_SEP.join(_inherit_col_sql(c) for c in (
|
||||
"user_id", "session_key", "chat_id", "chat_type", "thread_id", "display_name", "origin_json",
|
||||
))
|
||||
+ "\n WHERE id = ? AND parent_session_id IS NOT NULL\n"
|
||||
" AND EXISTS (\n"
|
||||
" SELECT 1 FROM sessions p\n"
|
||||
" WHERE p.id = sessions.parent_session_id\n"
|
||||
" AND p.end_reason = 'compression'\n"
|
||||
" )"
|
||||
)
|
||||
|
||||
|
||||
class SessionSessionsMixin:
|
||||
"""Session rows: create/inherit, lifecycle flags, model_config, listing, deletion."""
|
||||
|
||||
@@ -235,55 +256,8 @@ class SessionSessionsMixin:
|
||||
gateway re-records the peer would strand the child unroutable); delegate
|
||||
children must NOT inherit them (peer recovery could repoint gateway
|
||||
traffic into a subagent's session)."""
|
||||
conn.execute(
|
||||
f"""UPDATE sessions
|
||||
SET cwd = COALESCE(sessions.cwd,
|
||||
(SELECT p.cwd FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
git_repo_root = COALESCE(sessions.git_repo_root,
|
||||
(SELECT p.git_repo_root FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
git_branch = COALESCE(sessions.git_branch,
|
||||
(SELECT p.git_branch FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
profile_name = COALESCE(sessions.profile_name,
|
||||
(SELECT p.profile_name FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id
|
||||
AND ({_SAME_KEY_NAMESPACE_SQL})))
|
||||
WHERE id = ? AND parent_session_id IS NOT NULL""",
|
||||
(session_id,),
|
||||
)
|
||||
conn.execute(
|
||||
"""UPDATE sessions
|
||||
SET user_id = COALESCE(sessions.user_id,
|
||||
(SELECT p.user_id FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
session_key = COALESCE(sessions.session_key,
|
||||
(SELECT p.session_key FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
chat_id = COALESCE(sessions.chat_id,
|
||||
(SELECT p.chat_id FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
chat_type = COALESCE(sessions.chat_type,
|
||||
(SELECT p.chat_type FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
thread_id = COALESCE(sessions.thread_id,
|
||||
(SELECT p.thread_id FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
display_name = COALESCE(sessions.display_name,
|
||||
(SELECT p.display_name FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id)),
|
||||
origin_json = COALESCE(sessions.origin_json,
|
||||
(SELECT p.origin_json FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id))
|
||||
WHERE id = ? AND parent_session_id IS NOT NULL
|
||||
AND EXISTS (
|
||||
SELECT 1 FROM sessions p
|
||||
WHERE p.id = sessions.parent_session_id
|
||||
AND p.end_reason = 'compression'
|
||||
)""",
|
||||
(session_id,),
|
||||
)
|
||||
conn.execute(_INHERIT_PARENT_META_SQL, (session_id,))
|
||||
conn.execute(_INHERIT_PARENT_ROUTING_SQL, (session_id,))
|
||||
|
||||
def _insert_session_row(
|
||||
self, session_id: str, source: str, model: str = None, model_config: Dict[str, Any] = None,
|
||||
@@ -408,10 +382,8 @@ class SessionSessionsMixin:
|
||||
return str(exact[0]["id"])
|
||||
if len(rows) > 1:
|
||||
return None
|
||||
elif len(rows) > 1:
|
||||
distinct_users = {u for u in (str(r.get("user_id") or "").strip() for r in rows) if u}
|
||||
if len(distinct_users) > 1:
|
||||
return None
|
||||
elif len({u for u in (str(r.get("user_id") or "").strip() for r in rows) if u}) > 1:
|
||||
return None
|
||||
return str(rows[0]["id"])
|
||||
|
||||
# Orphaned gateway-session repair: widest plausible gap between a keyed
|
||||
@@ -433,35 +405,35 @@ class SessionSessionsMixin:
|
||||
"""Mark a session ended; the first end_reason wins (a compression split must
|
||||
keep ``'compression'`` even if a stale end_session() targets it later).
|
||||
reopen_session() first to deliberately re-end with a new reason."""
|
||||
def _do(conn):
|
||||
changed = conn.execute(
|
||||
"UPDATE sessions SET ended_at = ?, end_reason = ? "
|
||||
"WHERE id = ? AND ended_at IS NULL",
|
||||
(time.time(), end_reason, session_id),
|
||||
).rowcount
|
||||
# Only a boundary this call wrote advances the generation (a no-op must not rotate the peer).
|
||||
if changed:
|
||||
self._bump_conversation_generation(conn, session_id, end_reason)
|
||||
self._execute_write(_do)
|
||||
self._execute_write(lambda conn: self._end_and_bump(
|
||||
conn, "UPDATE sessions SET ended_at = ?, end_reason = ? WHERE id = ? AND ended_at IS NULL",
|
||||
(time.time(), end_reason, session_id), session_id, end_reason,
|
||||
))
|
||||
|
||||
def _end_and_bump(self, conn, sql: str, params: tuple, session_id: str, reason: str) -> int:
|
||||
"""Run an end-stamp UPDATE; only a boundary this call actually wrote advances the
|
||||
conversation generation (a no-op must not rotate the peer). Returns rowcount."""
|
||||
changed = conn.execute(sql, params).rowcount
|
||||
if changed:
|
||||
self._bump_conversation_generation(conn, session_id, reason)
|
||||
return changed
|
||||
|
||||
def reopen_session(self, session_id: str) -> None:
|
||||
"""Clear ended_at/end_reason so a session can be resumed; first stamp
|
||||
markerless legacy reset children that depend on the parent's mutable
|
||||
end_reason (WHERE shared with the listing predicate so they cannot drift)."""
|
||||
def _do(conn):
|
||||
placeholders = _session_ids_placeholders(_RESET_END_REASONS)
|
||||
conn.execute(
|
||||
"UPDATE sessions AS child SET model_config = json_set("
|
||||
"COALESCE(child.model_config, '{}'), '$._reset_from', child.parent_session_id) "
|
||||
"WHERE child.parent_session_id = ? "
|
||||
"AND json_extract(COALESCE(child.model_config, '{}'), "
|
||||
" '$._reset_from') IS NULL "
|
||||
f"AND {_legacy_reset_child_sql('child', placeholders)}",
|
||||
f"AND {_legacy_reset_child_sql('child', _session_ids_placeholders(_RESET_END_REASONS))}",
|
||||
(session_id, *_RESET_END_REASONS),
|
||||
)
|
||||
conn.execute(
|
||||
"UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?",
|
||||
(session_id,),
|
||||
"UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?", (session_id,),
|
||||
)
|
||||
self._execute_write(_do)
|
||||
|
||||
@@ -473,20 +445,15 @@ class SessionSessionsMixin:
|
||||
if not session_id:
|
||||
return False
|
||||
now = time.time()
|
||||
def _do(conn):
|
||||
cursor = conn.execute(
|
||||
"UPDATE sessions SET ended_at = ?, end_reason = ? "
|
||||
"WHERE id = ? AND (ended_at IS NULL "
|
||||
f"OR end_reason IN ({_RECOVERABLE_END_REASONS_SQL}))",
|
||||
(now, reason, session_id),
|
||||
)
|
||||
# /new and policy auto-resets promote rather than end_session, so the
|
||||
# generation advances here too — same transaction, only when written.
|
||||
if cursor.rowcount:
|
||||
self._bump_conversation_generation(conn, session_id, reason)
|
||||
return cursor.rowcount
|
||||
# /new and policy auto-resets promote rather than end_session, so the
|
||||
# generation advances here too — same transaction, only when written.
|
||||
try:
|
||||
return bool(self._execute_write(_do))
|
||||
return bool(self._execute_write(lambda conn: self._end_and_bump(
|
||||
conn,
|
||||
"UPDATE sessions SET ended_at = ?, end_reason = ? WHERE id = ? AND (ended_at IS NULL "
|
||||
f"OR end_reason IN ({_RECOVERABLE_END_REASONS_SQL}))",
|
||||
(now, reason, session_id), session_id, reason,
|
||||
)))
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
@@ -504,15 +471,12 @@ class SessionSessionsMixin:
|
||||
branch = (git_branch or "").strip()
|
||||
repo_root = (git_repo_root or "").strip()
|
||||
def _do(conn):
|
||||
current = conn.execute(
|
||||
"SELECT cwd FROM sessions WHERE id = ?", (session_id,)
|
||||
).fetchone()
|
||||
current = conn.execute("SELECT cwd FROM sessions WHERE id = ?", (session_id,)).fetchone()
|
||||
if current is None:
|
||||
return None
|
||||
current_cwd = current[0]
|
||||
sets = ["cwd = ?", "git_metadata_generation = COALESCE(git_metadata_generation, 0) + 1"]
|
||||
params: List[Any] = [cwd]
|
||||
if current_cwd != cwd or replace_git_meta:
|
||||
if current[0] != cwd or replace_git_meta:
|
||||
sets.extend(("git_branch = ?", "git_repo_root = ?"))
|
||||
params.extend((branch or None, repo_root or None))
|
||||
else: # same cwd: only overwrite with captured (non-empty) values
|
||||
@@ -520,8 +484,7 @@ class SessionSessionsMixin:
|
||||
if val:
|
||||
sets.append(f"{col} = ?")
|
||||
params.append(val)
|
||||
params.append(session_id)
|
||||
conn.execute(f"UPDATE sessions SET {', '.join(sets)} WHERE id = ?", params)
|
||||
conn.execute(f"UPDATE sessions SET {', '.join(sets)} WHERE id = ?", [*params, session_id])
|
||||
row = conn.execute(
|
||||
"SELECT git_metadata_generation FROM sessions WHERE id = ?", (session_id,),
|
||||
).fetchone()
|
||||
@@ -554,8 +517,7 @@ class SessionSessionsMixin:
|
||||
pairs = [(root, cwd) for cwd, root in cwd_to_root.items() if root and cwd]
|
||||
if pairs:
|
||||
self._write_sql(
|
||||
"UPDATE sessions SET git_repo_root = ? "
|
||||
"WHERE cwd = ? AND COALESCE(git_repo_root, '') = ''",
|
||||
"UPDATE sessions SET git_repo_root = ? WHERE cwd = ? AND COALESCE(git_repo_root, '') = ''",
|
||||
pairs, many=True,
|
||||
)
|
||||
|
||||
@@ -569,13 +531,14 @@ class SessionSessionsMixin:
|
||||
if not session_id:
|
||||
return
|
||||
when = float(ts if ts is not None else time.time())
|
||||
desc = bound_activity_description(description)
|
||||
prov = normalize_activity_provenance(provenance).value
|
||||
self._write_sql(
|
||||
"UPDATE sessions SET last_activity_at = ?, "
|
||||
"last_activity_description = ?, last_activity_provenance = ? "
|
||||
"WHERE id = ? AND (last_activity_at IS NULL OR last_activity_at < ?)",
|
||||
(when, desc, prov, session_id, when),
|
||||
(
|
||||
when, bound_activity_description(description),
|
||||
normalize_activity_provenance(provenance).value, session_id, when,
|
||||
),
|
||||
patience_s=self._ACTIVITY_WRITE_PATIENCE_S,
|
||||
)
|
||||
|
||||
@@ -594,10 +557,8 @@ class SessionSessionsMixin:
|
||||
if row is not None and not row[0] and (not row[1] or row[1] == ActivityProvenance.UNKNOWN.value):
|
||||
return
|
||||
self._write_sql(
|
||||
"UPDATE sessions SET last_activity_description = ?, "
|
||||
"last_activity_provenance = ? WHERE id = ?",
|
||||
("", ActivityProvenance.UNKNOWN.value, session_id),
|
||||
patience_s=self._ACTIVITY_WRITE_PATIENCE_S,
|
||||
"UPDATE sessions SET last_activity_description = ?, last_activity_provenance = ? WHERE id = ?",
|
||||
("", ActivityProvenance.UNKNOWN.value, session_id), patience_s=self._ACTIVITY_WRITE_PATIENCE_S,
|
||||
)
|
||||
|
||||
def get_session_activity(self, session_id: str) -> Optional[Dict[str, Any]]:
|
||||
@@ -605,11 +566,10 @@ class SessionSessionsMixin:
|
||||
row = self.get_session(session_id) if session_id else None
|
||||
if not row:
|
||||
return None
|
||||
return build_activity_snapshot(
|
||||
last_activity_at=row.get("last_activity_at"),
|
||||
last_activity_description=row.get("last_activity_description"),
|
||||
last_activity_provenance=row.get("last_activity_provenance"),
|
||||
)
|
||||
return build_activity_snapshot(**{
|
||||
key: row.get(key)
|
||||
for key in ("last_activity_at", "last_activity_description", "last_activity_provenance")
|
||||
})
|
||||
|
||||
def update_session_meta(
|
||||
self, session_id: str, model_config_json: str, model: Optional[str] = None,
|
||||
@@ -624,10 +584,9 @@ class SessionSessionsMixin:
|
||||
def update_system_prompt(self, session_id: str, system_prompt: Optional[str]) -> None:
|
||||
"""Store the full assembled system prompt snapshot."""
|
||||
def _do(conn):
|
||||
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
||||
conn.execute(
|
||||
"UPDATE sessions SET system_prompt_hash = ?, system_prompt = NULL WHERE id = ?",
|
||||
(system_prompt_hash, session_id),
|
||||
(self._store_system_prompt(conn, system_prompt), session_id),
|
||||
)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
self._execute_write(_do)
|
||||
@@ -658,23 +617,22 @@ class SessionSessionsMixin:
|
||||
"UPDATE sessions SET model = ?, model_config = ?, "
|
||||
"system_prompt = NULL, system_prompt_hash = NULL WHERE id = ?",
|
||||
lambda merged: (model, merged, session_id),
|
||||
clear_prompts=True,
|
||||
)
|
||||
|
||||
def _write_model_config_patch(
|
||||
self, session_id: str, patch: Dict[str, Any],
|
||||
sql: str = "UPDATE sessions SET model_config = ? WHERE id = ?",
|
||||
params: Optional[Callable[[Optional[str]], tuple]] = None, *, clear_prompts: bool = False,
|
||||
params: Optional[Callable[[Optional[str]], tuple]] = None,
|
||||
) -> None:
|
||||
"""Merge ``patch`` into model_config then run ``sql`` with ``params(merged)``
|
||||
(default: plain model_config UPDATE) in one write transaction; no-op when
|
||||
the row doesn't exist. ``clear_prompts`` also GCs unreferenced system_prompts."""
|
||||
in one write transaction; no-op when the row doesn't exist. A custom ``sql``
|
||||
(the prompt-nulling variants) also GCs unreferenced system_prompts."""
|
||||
def _do(conn):
|
||||
merged = self._merge_model_config_json(conn, session_id, patch)
|
||||
if merged is _MODEL_CONFIG_ROW_MISSING:
|
||||
return
|
||||
conn.execute(sql, params(merged) if params else (merged, session_id))
|
||||
if clear_prompts:
|
||||
if params is not None:
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
self._execute_write(_do)
|
||||
|
||||
@@ -685,9 +643,7 @@ class SessionSessionsMixin:
|
||||
that keeps ``_branched_from``/``_delegate_from`` alive); ``None`` deletes a
|
||||
key. Returns serialized JSON (``None`` when empty, matching create_session's
|
||||
NULL) or ``_MODEL_CONFIG_ROW_MISSING`` (``on_missing="raise"`` → ValueError)."""
|
||||
row = conn.execute(
|
||||
"SELECT model_config FROM sessions WHERE id = ?", (session_id,),
|
||||
).fetchone()
|
||||
row = conn.execute("SELECT model_config FROM sessions WHERE id = ?", (session_id,)).fetchone()
|
||||
if row is None:
|
||||
if on_missing == "raise":
|
||||
raise ValueError(f"Session not found: {session_id}")
|
||||
@@ -721,8 +677,7 @@ class SessionSessionsMixin:
|
||||
markers survive); null system_prompt so cached footers cannot lie."""
|
||||
lock = {
|
||||
"provider": provider or "", "model": model or "", "model_options": model_options or {},
|
||||
"route_source": route_source or "", "confirmed": bool(confirmed),
|
||||
"updated_at": time.time(),
|
||||
"route_source": route_source or "", "confirmed": bool(confirmed), "updated_at": time.time(),
|
||||
}
|
||||
self._write_model_config_patch(
|
||||
session_id, {"browser_model_lock": lock},
|
||||
@@ -733,7 +688,6 @@ class SessionSessionsMixin:
|
||||
system_prompt_hash = NULL
|
||||
WHERE id = ?""",
|
||||
lambda merged: (merged, model, session_id),
|
||||
clear_prompts=True,
|
||||
)
|
||||
|
||||
def set_session_yolo(self, session_id: str, enabled: bool) -> None:
|
||||
@@ -761,8 +715,7 @@ class SessionSessionsMixin:
|
||||
self.flush_token_counts()
|
||||
row = self._read_one(
|
||||
"SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved "
|
||||
"FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
||||
"WHERE s.id = ?",
|
||||
"FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash WHERE s.id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
return self._session_row_dict(row) if row else None
|
||||
@@ -793,12 +746,11 @@ class SessionSessionsMixin:
|
||||
exact = self.get_session(session_id_or_prefix)
|
||||
if exact:
|
||||
return exact["id"]
|
||||
escaped = _escape_like(session_id_or_prefix)
|
||||
matches = [row["id"] for row in self._read_all(
|
||||
matches = self._read_all(
|
||||
"SELECT id FROM sessions WHERE id LIKE ? ESCAPE '\\' ORDER BY started_at DESC LIMIT 2",
|
||||
(f"{escaped}%",),
|
||||
)]
|
||||
return matches[0] if len(matches) == 1 else None
|
||||
(f"{_escape_like(session_id_or_prefix)}%",),
|
||||
)
|
||||
return matches[0]["id"] if len(matches) == 1 else None
|
||||
|
||||
def backfill_null_session_profiles(self, profile_name: str) -> int:
|
||||
"""Stamp this store's own profile onto legacy ``profile_name IS NULL`` rows,
|
||||
@@ -853,7 +805,7 @@ class SessionSessionsMixin:
|
||||
|
||||
def set_session_archived(self, session_id: str, archived: bool) -> bool:
|
||||
"""Soft-hide (or unhide) a session and its compression lineage; messages are kept."""
|
||||
return self._set_lineage_column('archived', session_id, 1 if archived else 0)
|
||||
return self._set_lineage_column("archived", session_id, int(archived))
|
||||
|
||||
# Accidental end reasons recovery treats as resumable (also interpolated into
|
||||
# the recovery/promotion SQL so literals cannot drift).
|
||||
@@ -892,18 +844,18 @@ class SessionSessionsMixin:
|
||||
def set_session_pinned(self, session_id: str, pinned: bool) -> bool:
|
||||
"""Pin/unpin a session and its compression lineage (pins are exempt from the
|
||||
``sessions.auto_archive`` sweep)."""
|
||||
return self._set_lineage_column('pinned', session_id, 1 if pinned else 0)
|
||||
return self._set_lineage_column("pinned", session_id, int(pinned))
|
||||
|
||||
def set_session_hidden(self, session_id: str, hidden: bool) -> bool:
|
||||
"""Hide/unhide a session and its compression lineage from the default listing;
|
||||
it stays resumable by the owning surface."""
|
||||
return self._set_lineage_column('hidden', session_id, 1 if hidden else 0)
|
||||
return self._set_lineage_column("hidden", session_id, int(hidden))
|
||||
|
||||
def set_session_read(self, session_id: str, read: bool = True) -> bool:
|
||||
"""Mark read/unread across the compression lineage. ``last_read_at`` is a
|
||||
watermark: unread when activity postdates it (no write on the message
|
||||
path). NULL = never tracked = read; 0 = explicitly unread."""
|
||||
return self._set_lineage_column('last_read_at', session_id, time.time() if read else 0.0)
|
||||
return self._set_lineage_column("last_read_at", session_id, time.time() if read else 0.0)
|
||||
|
||||
@staticmethod
|
||||
def session_unread(session_row: Dict[str, Any]) -> bool:
|
||||
@@ -928,29 +880,28 @@ class SessionSessionsMixin:
|
||||
finds ``AN-94``); chain membership keeps the leading-wildcard LIKE bounded."""
|
||||
params: List[Any] = []
|
||||
clauses: List[str] = []
|
||||
def _like_pattern(needle: str) -> str:
|
||||
def like(needle: str) -> str:
|
||||
return f"%{_escape_like(needle)}%"
|
||||
if id_needle:
|
||||
clauses.append(
|
||||
"EXISTS (SELECT 1 FROM chain cq WHERE cq.root_id = s.id"
|
||||
" AND LOWER(cq.cur_id) LIKE ? ESCAPE '\\')"
|
||||
)
|
||||
params.append(_like_pattern(id_needle))
|
||||
params.append(like(id_needle))
|
||||
if search_needle:
|
||||
compact_needle = re.sub(r"[\W_]+", "", search_needle)
|
||||
compact_sql = (
|
||||
"REPLACE(REPLACE(REPLACE(REPLACE(LOWER(COALESCE({0}, '')),"
|
||||
" '-', ''), '_', ''), '.', ''), ' ', '')"
|
||||
)
|
||||
search_clause = (
|
||||
"EXISTS (SELECT 1 FROM chain cq JOIN sessions cs ON cs.id = cq.cur_id"
|
||||
" WHERE cq.root_id = s.id AND (LOWER(COALESCE(cs.title, '')) LIKE ? ESCAPE '\\'"
|
||||
" OR LOWER(cq.cur_id) LIKE ? ESCAPE '\\'"
|
||||
)
|
||||
params.extend([_like_pattern(search_needle)] * 2)
|
||||
params.extend([like(search_needle)] * 2)
|
||||
if compact_needle:
|
||||
search_clause += f" OR {compact_sql.format('cs.title')} LIKE ? ESCAPE '\\'"
|
||||
params.append(_like_pattern(compact_needle))
|
||||
search_clause += (
|
||||
" OR REPLACE(REPLACE(REPLACE(REPLACE(LOWER(COALESCE(cs.title, '')),"
|
||||
" '-', ''), '_', ''), '.', ''), ' ', '') LIKE ? ESCAPE '\\'"
|
||||
)
|
||||
params.append(like(compact_needle))
|
||||
clauses.append(search_clause + "))")
|
||||
if not clauses:
|
||||
return where_sql, params
|
||||
@@ -961,37 +912,33 @@ class SessionSessionsMixin:
|
||||
"""Replace each compression root's surfaced fields with its live tip's (root
|
||||
``started_at`` kept for stable ordering), one batched query. ``_lineage_ids``
|
||||
carries every id on the chain: a persisted tile can hold a MIDDLE segment's id."""
|
||||
tip_ids_by_root: Dict[str, str] = {}
|
||||
chain_by_root: Dict[str, List[str]] = {}
|
||||
chain_by_root: Dict[str, List[str]] = {} # only roots whose tip differs from themselves
|
||||
for s in sessions:
|
||||
if s.get("end_reason") != "compression":
|
||||
continue
|
||||
chain = self.get_compression_chain(s["id"])
|
||||
tip_id = chain[-1] if chain else s["id"]
|
||||
if tip_id != s["id"]:
|
||||
tip_ids_by_root[s["id"]] = tip_id
|
||||
chain_by_root[s["id"]] = chain
|
||||
if s.get("end_reason") == "compression":
|
||||
chain = self.get_compression_chain(s["id"])
|
||||
if chain and chain[-1] != s["id"]:
|
||||
chain_by_root[s["id"]] = chain
|
||||
tip_rows = (
|
||||
self._get_session_rich_rows_batch(set(tip_ids_by_root.values()), compact_rows=compact_rows)
|
||||
if tip_ids_by_root else {}
|
||||
self._get_session_rich_rows_batch(
|
||||
{chain[-1] for chain in chain_by_root.values()}, compact_rows=compact_rows,
|
||||
) if chain_by_root else {}
|
||||
)
|
||||
projected = []
|
||||
for s in sessions:
|
||||
tip_id = tip_ids_by_root.get(s["id"])
|
||||
tip_row = tip_rows.get(tip_id) if tip_id else None
|
||||
chain = chain_by_root.get(s["id"])
|
||||
tip_row = tip_rows.get(chain[-1]) if chain else None
|
||||
if not tip_row:
|
||||
projected.append(s)
|
||||
continue
|
||||
merged = dict(s)
|
||||
for key in (
|
||||
"id", "ended_at", "end_reason", "message_count",
|
||||
"tool_call_count", "title", "last_active", "preview",
|
||||
"model", "system_prompt", "cwd", "git_branch", "git_repo_root",
|
||||
"id", "ended_at", "end_reason", "message_count", "tool_call_count", "title", "last_active",
|
||||
"preview", "model", "system_prompt", "cwd", "git_branch", "git_repo_root",
|
||||
):
|
||||
if key in tip_row:
|
||||
merged[key] = tip_row[key]
|
||||
merged["_lineage_root_id"] = s["id"]
|
||||
merged["_lineage_ids"] = chain_by_root.get(s["id"]) or None
|
||||
merged["_lineage_ids"] = chain
|
||||
projected.append(merged)
|
||||
return projected
|
||||
|
||||
@@ -1027,23 +974,19 @@ class SessionSessionsMixin:
|
||||
where_clauses.append("s.hidden = 0")
|
||||
where_sql = _where_sql(where_clauses)
|
||||
base_where_params = list(params) # pinned back-fill reuses the WHERE before LIMIT/OFFSET
|
||||
prompt_select = (
|
||||
"" if compact_rows
|
||||
else ", COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved"
|
||||
)
|
||||
prompt_join = (
|
||||
"" if compact_rows
|
||||
else "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash"
|
||||
prompt_select, prompt_join = ("", "") if compact_rows else (
|
||||
", COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved",
|
||||
"LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash",
|
||||
)
|
||||
_sel = self._compact_session_cols() if compact_rows else "s.*"
|
||||
id_needle = (id_query or "").strip().lower()
|
||||
search_needle = (search_query or "").strip().lower()
|
||||
if order_by_last_active:
|
||||
# The CTE walks compression-continuation edges forward from the admitted
|
||||
# rows; MAX over the chain gives effective_last_active in SQL. Do NOT
|
||||
# require child.started_at >= parent.ended_at: races insert the
|
||||
# continuation before ended_at is written.
|
||||
outer_where, id_params = self._chain_search_where(where_sql, id_needle, search_needle)
|
||||
outer_where, id_params = self._chain_search_where(
|
||||
where_sql, (id_query or "").strip().lower(), (search_query or "").strip().lower(),
|
||||
)
|
||||
query = f"""
|
||||
WITH RECURSIVE chain(root_id, cur_id) AS (
|
||||
SELECT s.id, s.id FROM sessions s {where_sql}
|
||||
@@ -1093,7 +1036,7 @@ class SessionSessionsMixin:
|
||||
# projects to its tip like any other row.
|
||||
if include_pinned:
|
||||
seen_ids = {s["id"] for s in sessions}
|
||||
pinned_where = (f"{where_sql} AND s.pinned = 1" if where_sql else "WHERE s.pinned = 1")
|
||||
pinned_where = f"{where_sql} AND s.pinned = 1" if where_sql else "WHERE s.pinned = 1"
|
||||
pinned_query = f"""
|
||||
SELECT {_sel}{prompt_select},
|
||||
{_PREVIEW_COL_SQL},
|
||||
@@ -1125,8 +1068,7 @@ class SessionSessionsMixin:
|
||||
if not ids:
|
||||
return {}
|
||||
statuses: Dict[str, str] = {sid: "empty" for sid in ids}
|
||||
placeholders = _session_ids_placeholders(ids)
|
||||
query = f"""
|
||||
rows = self._read_all(f"""
|
||||
SELECT m.session_id, m.role,
|
||||
m.tool_calls IS NOT NULL AS has_tool_calls,
|
||||
m.finish_reason
|
||||
@@ -1134,11 +1076,10 @@ class SessionSessionsMixin:
|
||||
JOIN (
|
||||
SELECT session_id, MAX(id) AS max_id
|
||||
FROM messages
|
||||
WHERE session_id IN ({placeholders})
|
||||
WHERE session_id IN ({_session_ids_placeholders(ids)})
|
||||
GROUP BY session_id
|
||||
) latest ON m.id = latest.max_id
|
||||
"""
|
||||
rows = self._read_all(query, ids)
|
||||
""", ids)
|
||||
for row in rows:
|
||||
statuses[row["session_id"]] = classify_session_status(
|
||||
role=row["role"], has_tool_calls=bool(row["has_tool_calls"]),
|
||||
@@ -1158,8 +1099,7 @@ class SessionSessionsMixin:
|
||||
if max_messages == 0:
|
||||
return 0
|
||||
row = self._read_one(
|
||||
"SELECT COUNT(*) FROM ("
|
||||
"SELECT 1 FROM messages WHERE session_id = ? AND active = 1 LIMIT ?)",
|
||||
"SELECT COUNT(*) FROM (SELECT 1 FROM messages WHERE session_id = ? AND active = 1 LIMIT ?)",
|
||||
(session_id, max_messages + 1),
|
||||
)
|
||||
message_count = int(row[0] if row else 0)
|
||||
@@ -1173,21 +1113,15 @@ class SessionSessionsMixin:
|
||||
if not session_id:
|
||||
return False
|
||||
row = self._read_one("SELECT model_config FROM sessions WHERE id = ?", (session_id,))
|
||||
if row is None:
|
||||
return False
|
||||
return bool(_parse_model_config(row[0]).get("_branched_from"))
|
||||
return row is not None and bool(_parse_model_config(row[0]).get("_branched_from"))
|
||||
|
||||
def _session_lineage_root_to_tip(self, session_id: str) -> List[str]:
|
||||
if not session_id:
|
||||
return [session_id]
|
||||
chain = []
|
||||
chain: List[str] = []
|
||||
current = session_id
|
||||
seen = set()
|
||||
with self._read_ctx() as conn:
|
||||
for _ in range(100):
|
||||
if not current or current in seen:
|
||||
break
|
||||
seen.add(current)
|
||||
while current and current not in chain and len(chain) < 100:
|
||||
chain.append(current)
|
||||
row = conn.execute(
|
||||
"SELECT parent_session_id FROM sessions WHERE id = ?", (current,),
|
||||
@@ -1202,11 +1136,6 @@ class SessionSessionsMixin:
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Sessions MRU-first with a computed ``last_active``; ``workspace_key`` scopes
|
||||
to one workspace so ``hermes -c``/``--resume`` picks its last session."""
|
||||
select_with_last_active = (
|
||||
"SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved, "
|
||||
f"{_sql_session_last_active('s')} AS last_active "
|
||||
"FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
||||
)
|
||||
where_clauses = []
|
||||
params: list = []
|
||||
if source:
|
||||
@@ -1216,12 +1145,13 @@ class SessionSessionsMixin:
|
||||
ws_clause, ws_params = _workspace_key_clause(workspace_key)
|
||||
where_clauses.append(ws_clause)
|
||||
params.extend(ws_params)
|
||||
where_sql = _where_sql(where_clauses, " ")
|
||||
params.extend([limit, offset])
|
||||
return [self._session_row_dict(row) for row in self._read_all(
|
||||
f"{select_with_last_active}{where_sql} "
|
||||
"SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved, "
|
||||
f"{_sql_session_last_active('s')} AS last_active "
|
||||
"FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
||||
f"{_where_sql(where_clauses, ' ')} "
|
||||
"ORDER BY last_active DESC, s.started_at DESC, s.id DESC LIMIT ? OFFSET ?",
|
||||
params,
|
||||
[*params, limit, offset],
|
||||
)]
|
||||
|
||||
def session_count(
|
||||
@@ -1233,19 +1163,15 @@ class SessionSessionsMixin:
|
||||
total matches the listable rows."""
|
||||
where_clauses, params = _session_filter_where(
|
||||
exclude_children=exclude_children, source=source, sources=sources,
|
||||
exclude_sources=exclude_sources, cwd_prefix=cwd_prefix,
|
||||
min_message_count=min_message_count,
|
||||
exclude_sources=exclude_sources, cwd_prefix=cwd_prefix, min_message_count=min_message_count,
|
||||
archived_only=archived_only, include_archived=include_archived,
|
||||
)
|
||||
return self._read_one(
|
||||
f"SELECT COUNT(*) FROM sessions s{_where_sql(where_clauses, ' ')}", params,
|
||||
)[0]
|
||||
return self._read_one(f"SELECT COUNT(*) FROM sessions s{_where_sql(where_clauses, ' ')}", params)[0]
|
||||
|
||||
def session_count_ge(self, n: int = 1) -> bool:
|
||||
"""At least N sessions exist (archived included); LIMIT short-circuits
|
||||
instead of session_count()'s index scan."""
|
||||
rows = self._read_all("SELECT 1 FROM sessions LIMIT ?", (n,))
|
||||
return len(rows) >= n
|
||||
return len(self._read_all("SELECT 1 FROM sessions LIMIT ?", (n,))) >= n
|
||||
|
||||
def session_count_by_source(
|
||||
self, *, include_archived: bool = False, archived_only: bool = False,
|
||||
@@ -1253,16 +1179,15 @@ class SessionSessionsMixin:
|
||||
) -> Dict[str, int]:
|
||||
"""``{source: count}`` via one GROUP BY; ``exclude_children`` mirrors listing visibility."""
|
||||
where_clauses, params = _session_filter_where(
|
||||
exclude_children=exclude_children,
|
||||
archived_only=archived_only, include_archived=include_archived,
|
||||
exclude_children=exclude_children, archived_only=archived_only,
|
||||
include_archived=include_archived,
|
||||
)
|
||||
where_sql = _where_sql(where_clauses, " ")
|
||||
with self._read_ctx() as conn:
|
||||
if self._conn is None:
|
||||
raise RuntimeError("SessionDB connection is closed")
|
||||
rows = conn.execute(
|
||||
"SELECT COALESCE(NULLIF(s.source, ''), 'cli') AS source, COUNT(*) AS count "
|
||||
f"FROM sessions s{where_sql} "
|
||||
f"FROM sessions s{_where_sql(where_clauses, ' ')} "
|
||||
"GROUP BY COALESCE(NULLIF(s.source, ''), 'cli') ORDER BY count DESC",
|
||||
params,
|
||||
).fetchall()
|
||||
@@ -1274,7 +1199,7 @@ class SessionSessionsMixin:
|
||||
session = self.get_session(session_id)
|
||||
if not session:
|
||||
return False, ""
|
||||
return (self._is_explicit_fork_child_row(session), str(session.get("source") or "").strip())
|
||||
return self._is_explicit_fork_child_row(session), str(session.get("source") or "").strip()
|
||||
|
||||
@staticmethod
|
||||
def _remove_session_files(sessions_dir: Optional[Path], session_id: str) -> None:
|
||||
@@ -1310,7 +1235,7 @@ class SessionSessionsMixin:
|
||||
branch/compression children are orphaned. *expected_delete_ids*: proceed
|
||||
only if parent + delegate cascade still equals that set (re-walked inside
|
||||
the transaction on purpose: export-before-delete fails closed)."""
|
||||
removed_delegate_ids: List[str] = []
|
||||
removed_ids: List[str] = []
|
||||
expected_ids = set(expected_delete_ids) if expected_delete_ids is not None else None
|
||||
def _do(conn):
|
||||
if conn.execute("SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,)).fetchone() is None:
|
||||
@@ -1319,19 +1244,18 @@ class SessionSessionsMixin:
|
||||
session_id, *_collect_delegate_child_ids(conn, [session_id])
|
||||
}:
|
||||
return False
|
||||
removed_delegate_ids.extend(_delete_delegate_children(conn, [session_id]))
|
||||
removed_ids.extend(_delete_delegate_children(conn, [session_id]))
|
||||
conn.execute( # orphan remaining children (branches) so FK is satisfied
|
||||
"UPDATE sessions SET parent_session_id = NULL WHERE parent_session_id = ?",
|
||||
(session_id,),
|
||||
"UPDATE sessions SET parent_session_id = NULL WHERE parent_session_id = ?", (session_id,),
|
||||
)
|
||||
conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,))
|
||||
conn.execute("DELETE FROM sessions WHERE id = ?", (session_id,))
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
removed_ids.append(session_id)
|
||||
return True
|
||||
deleted = self._execute_write(_do)
|
||||
if deleted:
|
||||
for sid in removed_delegate_ids + [session_id]:
|
||||
self._remove_session_files(sessions_dir, sid)
|
||||
for sid in removed_ids:
|
||||
self._remove_session_files(sessions_dir, sid)
|
||||
return bool(deleted)
|
||||
|
||||
def delete_session_if_empty(self, session_id: str, sessions_dir: Optional[Path] = None) -> bool:
|
||||
@@ -1359,7 +1283,7 @@ class SessionSessionsMixin:
|
||||
deleted = self._execute_write(_do)
|
||||
if deleted:
|
||||
self._remove_session_files(sessions_dir, session_id)
|
||||
return bool(deleted)
|
||||
return deleted
|
||||
|
||||
def delete_sessions(self, session_ids: List[str], sessions_dir: Optional[Path] = None) -> int:
|
||||
"""Bulk delete with :meth:`delete_session` semantics per row, in ONE
|
||||
@@ -1376,17 +1300,13 @@ class SessionSessionsMixin:
|
||||
).fetchall()]
|
||||
if not existing:
|
||||
return 0
|
||||
existing_placeholders = _session_ids_placeholders(existing)
|
||||
ph = _session_ids_placeholders(existing)
|
||||
removed_ids.extend(_delete_delegate_children(conn, existing))
|
||||
conn.execute( # orphan children whose parent is in the kill list (FK)
|
||||
f"UPDATE sessions SET parent_session_id = NULL "
|
||||
f"WHERE parent_session_id IN ({existing_placeholders})",
|
||||
existing,
|
||||
f"UPDATE sessions SET parent_session_id = NULL WHERE parent_session_id IN ({ph})", existing,
|
||||
)
|
||||
conn.execute(
|
||||
f"DELETE FROM messages WHERE session_id IN ({existing_placeholders})", existing,
|
||||
)
|
||||
conn.execute(f"DELETE FROM sessions WHERE id IN ({existing_placeholders})", existing)
|
||||
conn.execute(f"DELETE FROM messages WHERE session_id IN ({ph})", existing)
|
||||
conn.execute(f"DELETE FROM sessions WHERE id IN ({ph})", existing)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
removed_ids.extend(existing)
|
||||
return len(existing)
|
||||
@@ -1419,9 +1339,8 @@ class SessionSessionsMixin:
|
||||
if not session_ids:
|
||||
return 0
|
||||
conn.execute(
|
||||
f"UPDATE sessions SET parent_session_id = NULL "
|
||||
f"WHERE parent_session_id IN ({_session_ids_placeholders(session_ids)})",
|
||||
list(session_ids),
|
||||
"UPDATE sessions SET parent_session_id = NULL "
|
||||
f"WHERE parent_session_id IN ({_session_ids_placeholders(session_ids)})", list(session_ids),
|
||||
)
|
||||
for sid in session_ids:
|
||||
# DELETE FROM messages: a row inserted between the SELECT and here
|
||||
@@ -1469,8 +1388,7 @@ class SessionSessionsMixin:
|
||||
self.set_meta("last_auto_archive", str(now))
|
||||
if archived > 0:
|
||||
logger.info(
|
||||
"state.db auto-archive: archived %d session(s) idle >= %s days", archived,
|
||||
idle_days,
|
||||
"state.db auto-archive: archived %d session(s) idle >= %s days", archived, idle_days,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("state.db auto-archive failed: %s", exc)
|
||||
|
||||
@@ -139,7 +139,7 @@ class TestConcurrentReadersDoNotRaceTheWriter:
|
||||
ALLOWED_FUNCS = {
|
||||
# Lifecycle: run before the instance is shared / after readers
|
||||
# are drained. Not reachable concurrently with writers.
|
||||
"__init__", "_connect_and_init",
|
||||
"__init__", "_open_writer", "_connect_and_init",
|
||||
"_connect_and_init_with_lock_patience", "close",
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user