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:
Teknium
2026-09-02 23:17:22 -07:00
parent c8616a9127
commit 54714ef066
3 changed files with 157 additions and 238 deletions

View File

@@ -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:

View File

@@ -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)

View File

@@ -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",
}