diff --git a/hermes_state_messages.py b/hermes_state_messages.py index fc42970b61..c6073d6585 100644 --- a/hermes_state_messages.py +++ b/hermes_state_messages.py @@ -734,28 +734,73 @@ class SessionMessagesMixin: missing = conn.execute(missing_sql, (session_id,)).fetchone() if missing is None: return True - rows = conn.execute( - "SELECT * FROM messages WHERE session_id = ? AND (active = 1 OR compacted = 1) ORDER BY id", - (session_id,)).fetchall() - first_id: Dict[Tuple[Any, ...], int] = {} - keyed_rows = [] - for row in rows: - key = self._display_dedupe_key(row) - first_id[key] = min(first_id.get(key, row["id"]), row["id"]) - keyed_rows.append((row, key)) - updates = [] - for row, key in keyed_rows: - order = first_id[key] - identity = self._display_identity(key) - if order != row["display_order"] or identity != row["display_identity"]: - updates.append((order, identity, row["id"])) - conn.executemany( - "UPDATE messages SET display_order = ?, display_identity = ? WHERE id = ?", - updates) + first_id: Dict[bytes, int] = {} + last_id = 0 + while True: + rows = conn.execute( + "SELECT id, role, content, timestamp, tool_call_id, tool_calls, tool_name, " + "display_kind, display_metadata, display_order, display_identity " + "FROM messages INDEXED BY idx_messages_session_id " + "WHERE session_id = ? AND id > ? AND (active = 1 OR compacted = 1) " + "ORDER BY id LIMIT 1000", + (session_id, last_id)) + batch_start = last_id + updates = [] + for row in rows: + last_id = row["id"] + identity = self._display_identity(self._display_dedupe_key(row)) + order = first_id.setdefault(identity, last_id) + if order != row["display_order"] or identity != row["display_identity"]: + updates.append((order, identity, last_id)) + rows.close() + if last_id == batch_start: + break + conn.executemany( + "UPDATE messages SET display_order = ?, display_identity = ? WHERE id = ?", updates) return True return bool(self._execute_write(_do)) + def _legacy_display_page(self, session_id: str, *, active_clause: str, limit: Optional[int], offset: int, + latest: bool) -> List[Any]: + """Project a legacy read-only display page without retaining transcript payloads.""" + representatives: Dict[bytes, Tuple[int, int]] = {} + with self._read_ctx() as conn: + conn.execute("BEGIN") + try: + has_session_index = conn.execute( + "SELECT 1 FROM sqlite_master WHERE type = 'index' AND name = ?", + ("idx_messages_session_id",), + ).fetchone() is not None + index_hint = "INDEXED BY idx_messages_session_id" if has_session_index else "NOT INDEXED" + rows = conn.execute( + "SELECT id, role, content, timestamp, tool_call_id, tool_calls, tool_name, active, " + f"display_kind, display_metadata FROM messages {index_hint} " + f"WHERE session_id = ?{active_clause} ORDER BY id ASC", + (session_id,)) + for row in rows: + identity = self._display_identity(self._display_dedupe_key(row)) + current = representatives.get(identity) + candidate = (row["active"], row["id"]) + if current is None or candidate > current: + representatives[identity] = candidate + rows.close() + + identities = list(representatives) + identities = identities[::-1][offset:][:limit][::-1] if latest else identities[offset:][:limit] + selected_ids = [representatives[identity][1] for identity in identities] + selected = {} + for start in range(0, len(selected_ids), 900): + chunk = selected_ids[start:start + 900] + selected.update({row["id"]: row for row in conn.execute( + f"SELECT * FROM messages WHERE session_id = ?{active_clause} " + f"AND id IN ({_placeholders(chunk)})", + (session_id, *chunk))}) + return [selected[row_id] for row_id in selected_ids if row_id in selected] + finally: + if conn.in_transaction: + conn.execute("ROLLBACK") + def _row_to_message_dict(self, row, *, warn_context: str, summary_flag: bool) -> Dict[str, Any]: """``dict(row)`` with content/tool_calls/display_metadata decoded; *summary_flag* keeps ``_compressed_summary`` only as ``True``.""" @@ -807,10 +852,10 @@ class SessionMessagesMixin: ORDER BY page.display_order ASC""" rows = self._read_all(sql, [session_id, -1 if limit is None else limit, offset, session_id]) elif include_compacted: - # Read-only legacy stores cannot persist display identities; retain the exact old projection. - rows = self._dedupe_display_generations(self._read_all( - "SELECT * FROM messages WHERE session_id = ?" + active_clause + " ORDER BY id ASC", [session_id])) - rows = rows[::-1][offset:][:limit][::-1] if latest else rows[offset:][:limit] + # Read-only legacy stores cannot persist display identities; keep only fixed-width + # identities and representative ids while scanning, then fetch the selected payloads. + rows = self._legacy_display_page( + session_id, active_clause=active_clause, limit=limit, offset=offset, latest=latest) else: sql = (f"SELECT * FROM messages WHERE session_id = ?{active_clause}" f"{' AND id > ?' if after_id is not None else ''} ORDER BY id {'DESC' if latest else 'ASC'}") diff --git a/tests/hermes_state/test_get_messages_include_compacted.py b/tests/hermes_state/test_get_messages_include_compacted.py index c61dc51f6c..cce2c40f64 100644 --- a/tests/hermes_state/test_get_messages_include_compacted.py +++ b/tests/hermes_state/test_get_messages_include_compacted.py @@ -14,6 +14,7 @@ job of ``include_inactive`` (audit / debug reads). import json import sqlite3 +import tracemalloc import pytest @@ -93,6 +94,27 @@ class TestIncludeCompacted: msgs = db.get_messages(sid, include_compacted=True) assert not any(not m["active"] and not m["compacted"] for m in msgs) + @pytest.mark.parametrize("read_only", [False, True]) + @pytest.mark.parametrize("latest", [False, True]) + @pytest.mark.parametrize("limit,offset", [(3, 1), (0, 0), (None, 2), (3, 100), (-1, 0), (3, -1)]) + def test_include_inactive_compacted_paging_keeps_rewound_rows(self, tmp_path, read_only, latest, limit, offset): + path = tmp_path / "audit.db" + db = _seed(SessionDB(path)) + expected = db.get_messages("s1", include_inactive=True) + db._execute_write(lambda conn: conn.execute( + "UPDATE messages SET display_order = NULL, display_identity = NULL")) + if read_only: + db.close() + db = SessionDB(path, read_only=True) + try: + all_rows = db.get_messages("s1", include_compacted=True, include_inactive=True) + assert all_rows == expected + page = db.get_messages("s1", include_compacted=True, include_inactive=True, + limit=limit, offset=offset, latest=latest) + assert page == (expected[::-1][offset:][:limit][::-1] if latest else expected[offset:][:limit]) + finally: + db.close() + def test_include_inactive_still_returns_everything(self, db): """Audit semantics are unchanged: include_inactive wins.""" sid = "s1" @@ -150,6 +172,99 @@ class TestDisplayDedupe: db._execute_write(_do) + @pytest.mark.parametrize("read_only", [False, True]) + def test_legacy_page_retains_only_bounded_payloads(self, tmp_path, read_only): + """Benjamin Brumbaugh's PR #106838: backfill and fallback retain identities, not payloads.""" + path = tmp_path / "legacy-large.db" + db = SessionDB(path) + sid = "legacy-large" + payload_size = 2_000_000 + db.create_session(sid, source="desktop") + db.append_messages_batch( + sid, [{"role": "assistant", "content": f"small-{i}"} for i in range(1_001)], + chunk_rows=500) + first_id = db.get_messages(sid, limit=1)[0]["id"] + self._copy_tail_as_new_generation(db, sid, [first_id]) + db.append_messages_batch(sid, [ + {"role": "assistant", "content": chr(65 + i) * payload_size} for i in range(12)]) + db._execute_write(lambda conn: conn.execute( + "UPDATE messages SET display_order = NULL, display_identity = NULL WHERE session_id = ?", (sid,))) + if read_only: + db.close() + db = SessionDB(path, read_only=True) + try: + tracemalloc.start() + try: + page = db.get_messages(sid, include_compacted=True, latest=True, limit=1) + _, peak = tracemalloc.get_traced_memory() + finally: + tracemalloc.stop() + assert page[0]["content"] == "L" * payload_size + assert peak < payload_size * 5 + if not read_only: + missing = db._read_one( + "SELECT COUNT(*) FROM messages WHERE display_order IS NULL OR display_identity IS NULL") + assert missing is not None and missing[0] == 0 + assert [row[0] for row in db._read_all( + "SELECT display_order FROM messages WHERE content = ? ORDER BY id", ("small-0",))] == [first_id, first_id] + finally: + db.close() + + def test_legacy_page_keeps_one_snapshot_during_rewind(self, tmp_path, monkeypatch): + """The identity scan and payload lookup cannot straddle a committed rewind.""" + path = tmp_path / "snapshot.db" + writer = SessionDB(path) + writer.create_session("snapshot", source="desktop") + writer.append_message("snapshot", "assistant", "visible-at-scan") + writer._execute_write(lambda conn: conn.execute( + "UPDATE messages SET display_order = NULL, display_identity = NULL")) + reader = SessionDB(path, read_only=True) + identity = reader._display_identity + + def rewind_after_identity(key): + result = identity(key) + writer._execute_write(lambda conn: conn.execute( + "UPDATE messages SET content = 'rewound', active = 0, compacted = 0")) + return result + + monkeypatch.setattr(reader, "_display_identity", rewind_after_identity) + try: + page = reader.get_messages("snapshot", include_compacted=True, latest=True, limit=1) + assert [(row["active"], row["content"]) for row in page] == [(1, "visible-at-scan")] + assert reader._conn is not None + assert not reader._conn.in_transaction + assert reader.get_messages("snapshot", include_compacted=True) == [] + finally: + reader.close() + writer.close() + + def test_legacy_page_preserves_aborted_transaction_error(self, tmp_path, monkeypatch): + path = tmp_path / "read-error.db" + writer = SessionDB(path) + writer.create_session("error", source="desktop") + writer.append_message("error", "assistant", "message") + writer._execute_write(lambda conn: conn.execute( + "UPDATE messages SET display_order = NULL, display_identity = NULL")) + writer.close() + reader = SessionDB(path, read_only=True) + identity = reader._display_identity + connection = reader._conn + assert connection is not None + + def abort_then_fail(key): + connection.execute("ROLLBACK") + raise sqlite3.OperationalError("disk I/O error") + + monkeypatch.setattr(reader, "_display_identity", abort_then_fail) + try: + with pytest.raises(sqlite3.OperationalError, match="disk I/O error"): + reader.get_messages("error", include_compacted=True, limit=1) + assert not connection.in_transaction + monkeypatch.setattr(reader, "_display_identity", identity) + assert reader.get_messages("error", include_compacted=True)[0]["content"] == "message" + finally: + reader.close() + def test_display_paging_and_append_work_is_bounded(self, db): """Page and identity-lookup work scale with the page, not the transcript: 10x rows must not cost 10x SQLite VM steps (the pre-index read deduped the whole session).""" @@ -215,7 +330,8 @@ class TestDisplayDedupe: large_steps = append_steps("write-large") assert large_steps < small_steps * 3 - def test_composite_handoff_keeps_live_turn_identity_and_first_position(self, db): + @pytest.mark.parametrize("read_only", [False, True]) + def test_composite_handoff_keeps_live_turn_identity_and_first_position(self, db, read_only): sid = "composite" db.create_session(sid, source="desktop") original_id = db.append_message(sid, role="user", content="live ask", timestamp=100.0) @@ -236,7 +352,12 @@ class TestDisplayDedupe: db._execute_write(lambda conn: conn.execute( "UPDATE messages SET display_order = NULL WHERE session_id = ?", (sid,))) - messages = db.get_messages(sid, include_compacted=True) + reader = SessionDB(db.db_path, read_only=True) if read_only else db + try: + messages = reader.get_messages(sid, include_compacted=True) + finally: + if read_only: + reader.close() assert [message["id"] for message in messages] == [carrier_id, later_id] assert messages[0]["content"] == carrier @@ -335,6 +456,7 @@ class TestDisplayDedupe: conn.execute("DROP INDEX IF EXISTS idx_messages_display_page") conn.execute("DROP INDEX IF EXISTS idx_messages_display_backfill") conn.execute("DROP INDEX IF EXISTS idx_messages_display_identity") + conn.execute("DROP INDEX IF EXISTS idx_messages_session_id") columns = {row[1] for row in conn.execute("PRAGMA table_info(messages)")} for column in ("display_order", "display_identity"): if column in columns: diff --git a/tests/tui_gateway/test_deferred_model_history.py b/tests/tui_gateway/test_deferred_model_history.py new file mode 100644 index 0000000000..85698130c7 --- /dev/null +++ b/tests/tui_gateway/test_deferred_model_history.py @@ -0,0 +1,155 @@ +"""Bounded Desktop hydration salvaged from Benjamin Brumbaugh's PR #106838.""" + +import threading + +import pytest + +from agent.replay_cleanup import canonicalize_replay_history +from hermes_state import SessionDB +from tui_gateway import server + + +@pytest.mark.parametrize("source,omit_messages", [("desktop", True), ("desktop", False), ("tui", True)]) +@pytest.mark.parametrize("profile", [None, "work"]) +def test_deferred_resume_preserves_model_history_and_db_ownership(tmp_path, monkeypatch, source, omit_messages, profile): + home = tmp_path / "work" + home.mkdir() + db = SessionDB(home / "state.db") + db.create_session("parent", source=source) + db.append_message("parent", "user", "ancestor-only display", timestamp=100.0) + db.end_session("parent", "compression") + db.create_session("tip", source=source, parent_session_id="parent") + db.append_message("tip", "user", "archived display", timestamp=101.0) + db.archive_and_compact("tip", [ + {"role": "assistant", "content": "summary", "_compressed_summary": True}, + {"role": "user", "content": "current ask"}, + {"role": "assistant", "content": "", "tool_calls": [ + {"id": "dangling", "type": "function", "function": {"name": "read_file", "arguments": "{}"}}]}, + ]) + expected = canonicalize_replay_history(db.get_messages_as_conversation( + "tip", repair_alternation=True, include_row_ids=True)) + stored = db.get_session("tip") + assert stored is not None + stored_count = stored["message_count"] + _, display = db.get_resume_conversations("tip") + prefix = db.get_ancestor_display_prefix("tip") + display_reads = [] + original_display = db.get_resume_conversations + original_prefix = db.get_ancestor_display_prefix + closed = threading.Event() + built = threading.Event() + events = [] + original_close = db.close + + def close(): + original_close() + closed.set() + + def read_display(sid): + display_reads.append("display") + return original_display(sid) + + def read_prefix(sid): + display_reads.append("prefix") + return original_prefix(sid) + + def acquire(db_path=None, **kwargs): + assert db_path == home / "state.db" + return db + + monkeypatch.setattr(db, "close", close) + monkeypatch.setattr(db, "get_resume_conversations", read_display) + monkeypatch.setattr(db, "get_ancestor_display_prefix", read_prefix) + monkeypatch.setattr("hermes_state_registry.acquire", acquire) + monkeypatch.setattr(server, "_profile_home", lambda p: home if p else None) + monkeypatch.setattr(server, "_profile_configured_cwd", lambda _: str(tmp_path)) + monkeypatch.setattr(server, "_default_session_cwd", lambda: str(tmp_path)) + monkeypatch.setattr(server, "_get_db", lambda: db) + monkeypatch.setattr(server, "_enable_gateway_prompts", lambda: None) + monkeypatch.setattr(server, "_schedule_session_cap_enforcement", lambda: None) + monkeypatch.setattr(server, "_maybe_schedule_auto_continue", lambda *args: None) + monkeypatch.setattr(server, "_start_agent_build", lambda *args: built.set()) + monkeypatch.setattr(server, "_emit", lambda kind, sid, payload: events.append((kind, payload))) + sid = None + try: + response = server.handle_request({"id": "resume", "method": "session.resume", "params": { + "session_id": "tip", "source": source, "defer_history": True, + "omit_messages": omit_messages, **({"profile": profile} if profile else {}), + }}) + assert response is not None and "error" not in response, response + sid = response["result"]["session_id"] + session = server._sessions[sid] + assert session["resume_history_ready"].wait(5) + assert built.wait(5) + model_only = source == "desktop" and omit_messages + assert display_reads == ([] if model_only else ["display", "prefix"]) + assert session["history"] == expected + assert session["display_history_prefix"] == ([] if model_only else prefix) + count = stored_count if model_only else len(display) + assert session["resume_message_count"] == count + assert ("session.resume_progress", {"message_count": count, "phase": "history", "status": "complete"}) in events + if profile: + assert closed.wait(5) + assert db._conn is None + else: + assert not closed.is_set() + assert db.get_session("tip") is not None + finally: + if sid is not None: + server._sessions.pop(sid, None) + db.close() + + +@pytest.mark.parametrize("outcome", ["replaced", "failed"]) +def test_model_hydration_discards_stale_results_and_closes_owned_db(tmp_path, monkeypatch, outcome): + db = SessionDB(tmp_path / "state.db") + db.create_session("stored", source="desktop") + db.append_message("stored", "user", "model context") + started, release, closed = threading.Event(), threading.Event(), threading.Event() + original_read, original_close = db.get_messages_as_conversation, db.close + events, builds = [], [] + old = {"history": [], "history_lock": threading.RLock(), "resume_hydrating": True, + "resume_history_ready": threading.Event(), "agent_ready": threading.Event(), + "resume_message_count": 1} + replacement = {"history": [{"role": "user", "content": "replacement"}]} + monkeypatch.setitem(server._sessions, "hydrating", old) + + def read(*args, **kwargs): + started.set() + assert release.wait(5) + if outcome == "failed": + db._read_all("SELECT * FROM deliberately_missing_table") + return original_read(*args, **kwargs) + + def close(): + original_close() + closed.set() + + monkeypatch.setattr(db, "get_messages_as_conversation", read) + monkeypatch.setattr(db, "close", close) + monkeypatch.setattr(server, "_emit", lambda *args: events.append(args)) + monkeypatch.setattr(server, "_start_agent_build", lambda *args: builds.append(args)) + try: + server._schedule_resume_hydration("hydrating", "stored", db, close_db=True, model_history_only=True) + assert started.wait(5) + assert not closed.is_set() + if outcome == "replaced": + server._sessions["hydrating"] = replacement + release.set() + assert closed.wait(5) + assert db._conn is None + assert builds == [] + if outcome == "replaced": + assert server._sessions["hydrating"] is replacement + assert replacement == {"history": [{"role": "user", "content": "replacement"}]} + assert old["history"] == [] + assert not any(payload.get("status") == "complete" for _, _, payload in events) + else: + assert "hydrating" not in server._sessions + assert old["resume_history_ready"].is_set() + assert "deliberately_missing_table" in old["resume_history_error"] + finally: + release.set() + assert closed.wait(5) + server._sessions.pop("hydrating", None) + original_close() diff --git a/tui_gateway/methods_session.py b/tui_gateway/methods_session.py index e15fcd3b87..01edabe3b0 100644 --- a/tui_gateway/methods_session.py +++ b/tui_gateway/methods_session.py @@ -752,7 +752,10 @@ def _resume_deferred(ctx: _Resume) -> dict: resume_message_count=int(ctx.found.get("message_count") or 0)) if (reused := ctx.claim(sid, record)) is not None: return reused - _schedule_resume_hydration(sid, ctx.target, ctx.db, close_db=ctx.owns_db) + # Desktop owns the visible transcript through bounded REST pages, not this model-history restore. + _schedule_resume_hydration( + sid, ctx.target, ctx.db, close_db=ctx.owns_db, + model_history_only=source == "desktop" and ctx.omit_messages) ctx.owns_db = False # the hydration worker now owns (and closes) the profile-scoped handle _schedule_session_cap_enforcement() return _resume_response(ctx, sid, record, info=ctx.info(cwd, overrides), messages=[], diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 2d05f45567..7f37b0f13a 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -2575,10 +2575,14 @@ def _schedule_agent_build(sid: str, delay: float = 0.05) -> None: timer.start() -def _load_resume_transcript(db, stored_id: str) -> tuple[list, list, list]: +def _load_resume_transcript(db, stored_id: str, *, model_history_only: bool = False) -> tuple[list, list, list]: """(raw_history, display_history, ancestor_prefix) for a cold resume. The full lineage is materialized only while it fits sessions.max_resume_messages (the transcript is REST-paginated), else the tip alone.""" from hermes_state import SessionResumeTooLargeError + if model_history_only: + raw_history = db.get_messages_as_conversation( + stored_id, repair_alternation=True, include_row_ids=True) + return raw_history, [], [] prefix_fits = True guard = getattr(db, "assert_resume_safe", None) if callable(guard): @@ -2597,7 +2601,8 @@ def _load_resume_transcript(db, stored_id: str) -> tuple[list, list, list]: return raw_history, raw_history, [] -def _schedule_resume_hydration(sid: str, stored_id: str, db, *, close_db: bool = False) -> None: +def _schedule_resume_hydration(sid: str, stored_id: str, db, *, close_db: bool = False, + model_history_only: bool = False) -> None: """Load a cold resume's transcript off the JSON-RPC response path.""" def _run() -> None: @@ -2607,22 +2612,24 @@ def _schedule_resume_hydration(sid: str, stored_id: str, db, *, close_db: bool = return _emit("session.resume_progress", sid, {"phase": "history", "status": "loading"}) db.reopen_session(stored_id) - raw_history, display_history, prefix = _load_resume_transcript(db, stored_id) + raw_history, display_history, prefix = _load_resume_transcript( + db, stored_id, model_history_only=model_history_only) # Display keeps the full transcript; the model-fed history uses the # same canonicalization as gateway resume and the send path. history = canonicalize_replay_history(raw_history) if _sessions.get(sid) is not session: return with session["history_lock"]: - session.update(history=history, display_history_prefix=prefix, resume_hydrating=False, - resume_message_count=len(display_history)) + session.update(history=history, display_history_prefix=prefix, resume_hydrating=False) + if not model_history_only: + session["resume_message_count"] = len(display_history) # Deferred resumes answered before the transcript existed; cache the derived todo snapshot now. todo_state = _todo_state_from_history(history) if todo_state is not None and session.get("todo_state") is None: session["todo_state"] = todo_state session["resume_history_ready"].set() _emit("session.resume_progress", sid, - {"message_count": len(display_history), "phase": "history", "status": "complete"}) + {"message_count": session["resume_message_count"], "phase": "history", "status": "complete"}) _maybe_schedule_auto_continue(sid, session, stored_id) _start_agent_build(sid, session) except Exception as exc: