diff --git a/hermes_state_compression.py b/hermes_state_compression.py index 124ffc7015..d6eae97865 100644 --- a/hermes_state_compression.py +++ b/hermes_state_compression.py @@ -236,11 +236,9 @@ class SessionCompressionMixin: _ceiling_clause = " AND id <= ?" _params.append(int(watermark_ceiling)) tail_ids, tail_tool_calls = self._tail_rows_after_watermark( - conn, - "SELECT id, tool_calls FROM messages " + conn, "SELECT id, tool_calls FROM messages " "WHERE session_id = ? AND active = 1 AND id > ?" - f"{_ceiling_clause} ORDER BY id", - _params, + f"{_ceiling_clause} ORDER BY id", _params, ) if tail_ids: self._clone_message_rows(conn, tail_ids, session_id=child_session_id) @@ -252,8 +250,7 @@ class SessionCompressionMixin: ) updated = conn.execute( "UPDATE sessions SET ended_at = ?, end_reason = 'compression' " - "WHERE id = ? AND ended_at IS NULL", - (time.time(), parent_session_id), + "WHERE id = ? AND ended_at IS NULL", (time.time(), parent_session_id), ) if updated.rowcount != 1: raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}") @@ -322,8 +319,7 @@ class SessionCompressionMixin: def _do(conn): cursor = conn.execute( "UPDATE sessions SET compression_failure_cooldown_until = ?, " - "compression_failure_error = ? WHERE id = ?", - (deadline, error, session_id), + "compression_failure_error = ? WHERE id = ?", (deadline, error, session_id), ) if cursor.rowcount != 1: raise RuntimeError(f"compression cooldown rollback session missing: {session_id}") @@ -348,8 +344,7 @@ class SessionCompressionMixin: self._write_sql_logged( "clear_compression_failure_cooldown", session_id, "UPDATE sessions SET compression_failure_cooldown_until = NULL, " - "compression_failure_error = NULL WHERE id = ?", - (session_id,), + "compression_failure_error = NULL WHERE id = ?", (session_id,), ) def _read_session_number(self, column: str, session_id: str, cast: type, zero: Any) -> Any: @@ -484,8 +479,7 @@ class SessionCompressionMixin: def _row(sid: str): row = conn.execute( "SELECT id, parent_session_id, source, model_config, end_reason " - "FROM sessions WHERE id = ?", - (sid,), + "FROM sessions WHERE id = ?", (sid,), ).fetchone() return dict(row) if row else None @@ -590,8 +584,7 @@ class SessionCompressionMixin: conversation_id = self._session_turn_lease_key_on_conn(conn, session_id) cursor = conn.execute( "UPDATE session_turn_leases SET expires_at = ? " - "WHERE conversation_id = ? AND holder = ?", - (expires_at, conversation_id, holder), + "WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder), ) return cursor.rowcount > 0 diff --git a/hermes_state_messages.py b/hermes_state_messages.py index a318c4487a..0d801db935 100644 --- a/hermes_state_messages.py +++ b/hermes_state_messages.py @@ -48,15 +48,18 @@ def _json_or(raw: Any, fallback: Any, warning: str) -> Any: return fallback -def _tool_calls_len(raw: Any) -> int: - """Tool-call count of a stored ``tool_calls`` column (0 unless a JSON list).""" +def _tool_calls_len(raw: Any, scalar: int = 0) -> int: + """Tool-call count of a stored ``tool_calls`` column: list length, else *scalar* + for a truthy non-list value, 0 for empty/undecodable.""" if not raw: return 0 try: parsed = json.loads(raw) if isinstance(raw, str) else raw except (TypeError, ValueError): return 0 - return len(parsed) if isinstance(parsed, list) else 0 + if isinstance(parsed, list): + return len(parsed) + return scalar if parsed else 0 def _coerce_timestamp(value: Any, default: float) -> float: @@ -91,6 +94,14 @@ def _ended_by_compression(row) -> bool: return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression" +# _rows_to_conversation copies these row columns verbatim when truthy, in this order. +# ``api_content`` is returned VERBATIM (no sanitize/strip): the replay path substitutes +# it to keep the provider prompt cache byte-stable. +_VERBATIM_COLS = ("api_content", "display_kind") +_META_COLS = ("timestamp", "tool_call_id", "tool_name", "effect_disposition") +_ASSISTANT_JSON_COLS = ("reasoning_details", "codex_reasoning_items", "codex_message_items") + + def _stale_holder(row, now: float) -> bool: """A lock/lease row whose holder is expired or a provably dead local process.""" from hermes_state import _compression_lock_holder_process_is_dead @@ -162,20 +173,19 @@ class SessionMessagesMixin: return None if isinstance(display_metadata, str): try: - parsed = json.loads(display_metadata) + display_metadata = json.loads(display_metadata) except (json.JSONDecodeError, TypeError): logger.warning("Ignoring non-JSON display metadata on write") return None - if not isinstance(parsed, dict): + if not isinstance(display_metadata, dict): logger.warning("Ignoring non-object display metadata on write") return None - return json.dumps(parsed) - if isinstance(display_metadata, dict): - return json.dumps(display_metadata) - logger.warning( - "Ignoring unexpected display metadata type on write: %s", type(display_metadata).__name__, - ) - return None + elif not isinstance(display_metadata, dict): + logger.warning( + "Ignoring unexpected display metadata type on write: %s", type(display_metadata).__name__, + ) + return None + return json.dumps(display_metadata) @staticmethod def _decode_display_metadata(raw: Any) -> Optional[Dict[str, Any]]: @@ -689,15 +699,13 @@ class SessionMessagesMixin: conn.execute( f"INSERT INTO messages ({col_list}, active, compacted) " f"SELECT {col_list}, 1, 0 FROM messages " - f"WHERE id IN ({placeholders}) ORDER BY id", - tail_ids, + f"WHERE id IN ({placeholders}) ORDER BY id", tail_ids, ) else: conn.execute( f"INSERT INTO messages ({col_list}, session_id, active, compacted) " f"SELECT {col_list}, ?, 1, 0 FROM messages " - f"WHERE id IN ({placeholders}) ORDER BY id", - [session_id, *tail_ids], + f"WHERE id IN ({placeholders}) ORDER BY id", [session_id, *tail_ids], ) def archive_and_compact( @@ -752,20 +760,18 @@ class SessionMessagesMixin: tail_tool_calls = 0 if watermark is not None: tail_ids, tail_tool_calls = self._tail_rows_after_watermark( - conn, - "SELECT id, tool_calls FROM messages " + conn, "SELECT id, tool_calls FROM messages " "WHERE session_id = ? AND active = 1 AND id > ? ORDER BY id", (session_id, int(watermark)), ) # Rewind targets sit AT/BELOW the watermark (the compressor only saw rows up # to it); without the bound a concurrent append would steal a LIMIT slot. - rewind_tail_ids: Optional[list[int]] = None + rewind_ids: list[int] = [] if tail_count > 0: if watermark is not None: tail_rows = conn.execute( "SELECT id FROM messages WHERE session_id = ? AND active = 1 AND id <= ? " - "ORDER BY id DESC LIMIT ?", - (session_id, int(watermark), int(tail_count)), + "ORDER BY id DESC LIMIT ?", (session_id, int(watermark), int(tail_count)), ).fetchall() else: tail_rows = conn.execute( @@ -773,26 +779,23 @@ class SessionMessagesMixin: "WHERE session_id = ? AND active = 1 ORDER BY id DESC LIMIT ?", (session_id, int(tail_count)), ).fetchall() - rewind_tail_ids = [int(row["id"]) for row in tail_rows] - rewind_ids = [*(rewind_tail_ids or []), *tail_ids] + rewind_ids = [int(row["id"]) for row in tail_rows] + rewind_ids += tail_ids if rewind_ids: placeholders = _placeholders(rewind_ids) conn.execute( "UPDATE messages SET active = 0, compacted = 0 " - f"WHERE session_id = ? AND id IN ({placeholders})", - [session_id, *rewind_ids], + f"WHERE session_id = ? AND id IN ({placeholders})", [session_id, *rewind_ids], ) conn.execute( "UPDATE messages SET active = 0, compacted = 1 " "WHERE session_id = ? AND active = 1 " - f"AND id NOT IN ({placeholders})", - [session_id, *rewind_ids], + f"AND id NOT IN ({placeholders})", [session_id, *rewind_ids], ) else: conn.execute( "UPDATE messages SET active = 0, compacted = 1 " - "WHERE session_id = ? AND active = 1", - (session_id,), + "WHERE session_id = ? AND active = 1", (session_id,), ) inserted, tool_calls_total = self._insert_message_rows(conn, session_id, compacted_messages) if tail_ids: @@ -959,7 +962,7 @@ class SessionMessagesMixin: chunk = ids[start : start + 900] rows = self._read_all( f"""SELECT session_id, content FROM messages - WHERE session_id IN ({",".join("?" * len(chunk))}) + WHERE session_id IN ({_placeholders(chunk)}) AND role = 'tool' AND content LIKE '%/pull/%' ORDER BY id ASC""", chunk, @@ -991,14 +994,9 @@ class SessionMessagesMixin: (session_id, around_message_id, window), ).fetchall() rows = list(reversed(before_rows)) + list(after_rows) - return { - "window": [ - self._row_to_message_dict(row, warn_context="get_messages_around", summary_flag=False) - for row in rows - ], - "messages_before": max(0, len(before_rows) - 1), # before_rows includes the anchor - "messages_after": len(after_rows), - } + window_msgs = [self._row_to_message_dict(r, warn_context="get_messages_around", summary_flag=False) for r in rows] + # before_rows includes the anchor itself. + return {"window": window_msgs, "messages_before": max(0, len(before_rows) - 1), "messages_after": len(after_rows)} def resolve_resume_session_id(self, session_id: str) -> str: """Redirect a resume target to the descendant session that holds the messages. @@ -1025,14 +1023,10 @@ class SessionMessagesMixin: best = None # deepest node with messages for _ in range(32): try: - row = conn.execute( + if conn.execute( "SELECT 1 FROM messages WHERE session_id = ? LIMIT 1", (current,), - ).fetchone() - except Exception: - return session_id - if row is not None: - best = current - try: + ).fetchone() is not None: + best = current child_row = conn.execute( "SELECT id FROM sessions AS child WHERE child.parent_session_id = ? " " AND json_extract(COALESCE(child.model_config, '{}'), '$._branched_from') IS NULL " @@ -1040,8 +1034,7 @@ class SessionMessagesMixin: " AND json_extract(COALESCE(child.model_config, '{}'), '$._reset_from') IS NULL " f" AND NOT {_legacy_reset_child_sql('child', _RESET_END_REASONS_SQL)} " " AND COALESCE(child.source, '') != 'tool' " - "ORDER BY child.started_at DESC, child.id DESC LIMIT 1", - (current,), + "ORDER BY child.started_at DESC, child.id DESC LIMIT 1", (current,), ).fetchone() except Exception: return session_id @@ -1062,8 +1055,7 @@ class SessionMessagesMixin: return conn.execute( f"{prefix}{self._CONVERSATION_ROW_COLUMNS} " f"FROM messages WHERE session_id IN ({_placeholders(session_ids)})" - f"{active_clause} ORDER BY id", - tuple(session_ids), + f"{active_clause} ORDER BY id", tuple(session_ids), ).fetchall() def get_messages_as_conversation( @@ -1129,9 +1121,8 @@ class SessionMessagesMixin: Every dict is stamped ``_DB_PERSISTED_MARKER_KEY`` at the source (born durable) so an identity-losing handoff never re-appends the whole transcript on flush. - ``_row_id`` is opt-in (gateway reactions). ``api_content`` is returned VERBATIM - (no sanitize/strip): the replay path substitutes it to keep the provider prompt - cache byte-stable. Reasoning fields are restored on assistant rows only. + ``_row_id`` is opt-in (gateway reactions). Reasoning fields are restored on + assistant rows only. Key order of each dict is stable (see the column tables). """ from hermes_state import _strip_background_review_harness, _strip_stale_tool_call_markers messages = [] @@ -1140,22 +1131,17 @@ class SessionMessagesMixin: content = self._decode_content(row["content"]) if row["role"] in {"user", "assistant"} and isinstance(content, str): content = sanitize_context(content).strip() - msg = {"role": row["role"], "content": content} - msg[_DB_PERSISTED_MARKER_KEY] = True + msg = {"role": row["role"], "content": content, _DB_PERSISTED_MARKER_KEY: True} if include_row_ids and row["id"] is not None: msg["_row_id"] = row["id"] - for col in ("api_content", "display_kind"): - if row[col]: - msg[col] = row[col] + msg.update((col, row[col]) for col in _VERBATIM_COLS if row[col]) if row["display_metadata"]: decoded = self._decode_display_metadata(row["display_metadata"]) if decoded is not None: msg["display_metadata"] = decoded if include_summary_markers and row["_compressed_summary"]: msg["_compressed_summary"] = True - for col in ("timestamp", "tool_call_id", "tool_name", "effect_disposition"): - if row[col]: - msg[col] = row[col] + msg.update((col, row[col]) for col in _META_COLS if row[col]) if row["tool_calls"]: msg["tool_calls"] = _json_or( row["tool_calls"], [], @@ -1167,16 +1153,13 @@ class SessionMessagesMixin: if row["observed"]: msg["observed"] = True if row["role"] == "assistant": - for col in ("finish_reason", "reasoning"): - if row[col]: - msg[col] = row[col] + msg.update((col, row[col]) for col in ("finish_reason", "reasoning") if row[col]) if row["reasoning_content"] is not None: msg["reasoning_content"] = row["reasoning_content"] - for col in ("reasoning_details", "codex_reasoning_items", "codex_message_items"): - if row[col]: - msg[col] = _json_or( - row[col], None, f"Failed to deserialize {col}, falling back to None", - ) + msg.update( + (col, _json_or(row[col], None, f"Failed to deserialize {col}, falling back to None")) + for col in _ASSISTANT_JSON_COLS if row[col] + ) exact_clone_key = None if include_ancestors: skip, exact_clone_key = self._dedupe_replayed_user(messages, msg, exact_user_clones) @@ -1198,9 +1181,7 @@ class SessionMessagesMixin: logger.info( "Repaired %d message-alternation violation(s) while " "restoring session %s — durable transcript kept them, " - "see repair_message_sequence", - repaired, - session_id, + "see repair_message_sequence", repaired, session_id, ) return messages @@ -1246,8 +1227,7 @@ class SessionMessagesMixin: """Count the rows a resume would materialize (see ``_resume_count_scope``).""" session_ids, active_clause = self._resume_count_scope(session_id, tip_only) row = self._read_one( - f"SELECT COUNT(*) FROM messages " - f"WHERE session_id IN ({_placeholders(session_ids)}) AND {active_clause}", + f"SELECT COUNT(*) FROM messages WHERE session_id IN ({_placeholders(session_ids)}) AND {active_clause}", tuple(session_ids), ) return int(row[0] if row else 0) @@ -1273,8 +1253,7 @@ class SessionMessagesMixin: "SELECT COUNT(*) FROM (" f"SELECT 1 FROM messages WHERE session_id IN ({_placeholders(session_ids)}) " f"AND {active_clause} LIMIT ?" - ")", - (*session_ids, max_messages + 1), + ")", (*session_ids, max_messages + 1), ) message_count = int(row[0] if row else 0) if message_count > max_messages: @@ -1303,13 +1282,10 @@ class SessionMessagesMixin: lineage = self._rows_to_conversation( rows, session_id=session_id, include_ancestors=True, repair_alternation=False, include_row_ids=True, ) - prefix: List[Dict[str, Any]] = [] - for message in lineage: - if message.get("_row_id") in ancestor_ids: - projected = message.copy() - projected.pop("_row_id", None) - prefix.append(projected) - return prefix + return [ + {k: v for k, v in message.items() if k != "_row_id"} + for message in lineage if message.get("_row_id") in ancestor_ids + ] def get_conversation_root(self, session_id: str) -> str: """ROOT id of *session_id*'s lineage — the stable conversation id across @@ -1384,20 +1360,7 @@ class SessionMessagesMixin: rows = conn.execute( "SELECT tool_calls FROM messages WHERE session_id = ? AND active = 1", (session_id,), ).fetchall() - tool_call_count = 0 - for row in rows: - raw = row[0] - if not raw: - continue - try: - decoded = json.loads(raw) if isinstance(raw, str) else raw - except (json.JSONDecodeError, TypeError): - continue - if isinstance(decoded, list): - tool_call_count += len(decoded) - elif decoded: - tool_call_count += 1 - return len(rows), tool_call_count + return len(rows), sum(_tool_calls_len(row[0], scalar=1) for row in rows) def _split_rewind_target(self, target_row: Dict[str, Any], expected_target_content: Any, preserve_compaction_handoff: bool): """Validate an active rewind target and return its handoff scaffold (or None). @@ -1453,7 +1416,7 @@ class SessionMessagesMixin: active_rows = conn.execute( "SELECT id FROM messages WHERE session_id = ? AND active = 1 ORDER BY id", (session_id,), ).fetchall() - if [int(active_row[0]) for active_row in active_rows] != expected_active_ids: + if [int(r[0]) for r in active_rows] != expected_active_ids: raise RuntimeError("active transcript changed before the rewind could be persisted") row = conn.execute( "SELECT * FROM messages WHERE id = ? AND session_id = ?", (target_message_id, session_id), @@ -1466,15 +1429,13 @@ class SessionMessagesMixin: f"rewind target must be a 'user' message (got role=" f"{target_row.get('role')!r}, id={target_message_id})" ) - replacement_message_id: Optional[int] = None - replacement: Optional[Dict[str, Any]] = None + replacement_message_id = replacement = None if preserve_compaction_handoff or expected_target_content is not None: replacement = self._split_rewind_target(target_row, expected_target_content, preserve_compaction_handoff) - cursor = conn.execute( + ids = [r[0] for r in conn.execute( "SELECT id FROM messages WHERE session_id = ? AND id >= ? AND active = 1", (session_id, target_message_id), - ) - ids = [r[0] for r in cursor.fetchall()] + ).fetchall()] if ids: conn.execute(f"UPDATE messages SET active = 0 WHERE id IN ({_placeholders(ids)})", ids) if replacement is not None: @@ -1491,8 +1452,7 @@ class SessionMessagesMixin: head_row = conn.execute( "SELECT MAX(id) FROM messages WHERE session_id = ? AND active = 1", (session_id,), ).fetchone() - new_head_id = head_row[0] if head_row and head_row[0] is not None else None - return target_row, ids, new_head_id, replacement_message_id + return target_row, ids, head_row[0] if head_row else None, replacement_message_id target_row, rewound, new_head_id, replacement_message_id = self._execute_write(_do) # Decode for the prompt-buffer prefill without a second fallible DB operation. @@ -1504,12 +1464,9 @@ class SessionMessagesMixin: def message_count(self, session_id: str = None) -> int: """Count messages, optionally for a specific session.""" - with self._read_ctx() as conn: - if session_id: - cursor = conn.execute("SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,)) - else: - cursor = conn.execute("SELECT COUNT(*) FROM messages") - return cursor.fetchone()[0] + if session_id: + return self._read_one("SELECT COUNT(*) FROM messages WHERE session_id = ?", (session_id,))[0] + return self._read_one("SELECT COUNT(*) FROM messages")[0] def has_platform_message_id(self, session_id: str, platform_message_id: str) -> bool: """True when a message with *platform_message_id* exists (partial index lookup; @@ -1538,12 +1495,9 @@ class SessionMessagesMixin: return False if not isinstance(cfg, dict): return False + markers = (cfg.get("_branched_from"), cfg.get("_delegate_from")) parent_id = session.get("parent_session_id") - branched = cfg.get("_branched_from") - delegated = cfg.get("_delegate_from") - if parent_id: - return branched == parent_id or delegated == parent_id - return branched is not None or delegated is not None + return parent_id in markers if parent_id else any(m is not None for m in markers) def is_explicit_fork_child(self, session_id: str) -> bool: """Public read-only view of :meth:`_is_explicit_fork_child_row`; a missing row @@ -1568,10 +1522,9 @@ class SessionMessagesMixin: "SELECT generation FROM conversation_generations WHERE source = ? AND session_key = ?", (source, session_key), ) - if row is None or row["generation"] is None: + if row is None or row["generation"] is None or int(row["generation"]) <= 0: return None - generation = int(row["generation"]) - return generation if generation > 0 else None + return int(row["generation"]) def clear_messages(self, session_id: str) -> None: """Delete all messages for a session and reset its counters.""" @@ -1625,7 +1578,7 @@ class SessionMessagesMixin: def _do(conn): ids = _find_affected(conn) if ids: - conn.execute(f"UPDATE messages SET content = '' WHERE id IN ({','.join('?' * len(ids))})", ids) + conn.execute(f"UPDATE messages SET content = '' WHERE id IN ({_placeholders(ids)})", ids) return ids affected_ids = self._execute_write(_do) diff --git a/hermes_state_titles.py b/hermes_state_titles.py index c6ac1dfdc6..64272588bf 100644 --- a/hermes_state_titles.py +++ b/hermes_state_titles.py @@ -175,8 +175,7 @@ class SessionTitlesMixin: 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.title = ?", - (title,), + "WHERE s.title = ?", (title,), ) return self._session_row_dict(row) if row else None diff --git a/hermes_state_usage.py b/hermes_state_usage.py index f84455053e..0b99762758 100644 --- a/hermes_state_usage.py +++ b/hermes_state_usage.py @@ -245,8 +245,7 @@ class SessionUsageMixin: # Writer stuck mid-apply: leave deltas unapplied rather than race it. logger.warning( "async token accounting: writer did not stop within %.0fs; " - "%d queued delta(s) not persisted", - join_timeout, len(self._token_queue), + "%d queued delta(s) not persisted", join_timeout, len(self._token_queue), ) return # Writer gone: apply leftovers synchronously under the same busy protocol. Wait @@ -366,8 +365,7 @@ class SessionUsageMixin: """ row = conn.execute( "SELECT model, billing_provider, billing_base_url, billing_mode " - "FROM sessions WHERE id = ?", - (session_id,), + "FROM sessions WHERE id = ?", (session_id,), ).fetchone() sess = dict(row) if (row is not None and not task) else {} eff_model = model or sess.get("model") or "unknown"