diff --git a/hermes_state.py b/hermes_state.py index 31b2256412..dc79f35f8d 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -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: diff --git a/hermes_state_sessions.py b/hermes_state_sessions.py index 36c9a7e6da..2efc8e2cff 100644 --- a/hermes_state_sessions.py +++ b/hermes_state_sessions.py @@ -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) diff --git a/tests/state/test_writer_conn_thread_safety.py b/tests/state/test_writer_conn_thread_safety.py index 6e460eb8f5..ab19a82663 100644 --- a/tests/state/test_writer_conn_thread_safety.py +++ b/tests/state/test_writer_conn_thread_safety.py @@ -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", }