From fd41164861575f5564cdb091dc16204d0f49883b Mon Sep 17 00:00:00 2001 From: poisdahl <4091911+poisdahl@users.noreply.github.com> Date: Sat, 22 Aug 2026 17:30:35 +0200 Subject: [PATCH] fix(history): keep carrier rewinds race-safe after refresh --- cli.py | 13 ++-- gateway/session.py | 16 +++-- gateway/slash_commands.py | 13 ++-- hermes_cli/web_routers/sessions.py | 11 ++-- hermes_state.py | 16 +++++ tests/gateway/test_retry_replacement.py | 59 ++++++++++++++++++ tests/gateway/test_undo_rewind_session.py | 61 +++++++++++++++++++ tests/hermes_cli/test_web_server.py | 17 +++++- .../test_composite_carrier_rewind.py | 30 +++++++++ tui_gateway/methods_session.py | 2 +- tui_gateway/methods_tools.py | 8 ++- tui_gateway/server.py | 30 +++------ 12 files changed, 226 insertions(+), 50 deletions(-) diff --git a/cli.py b/cli.py index 0d4a5a7386..746773dd47 100644 --- a/cli.py +++ b/cli.py @@ -10183,6 +10183,9 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): return sanitize_context(content).strip() return content + expected_active_ids = self._session_db.get_active_message_ids( + self.session_id + ) durable = self._session_db.get_messages_as_conversation( self.session_id, include_row_ids=True, @@ -10226,12 +10229,6 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): target_row_id = durable_target.get("_row_id") if not isinstance(target_row_id, int): raise RuntimeError("persisted rewind target has no row identity") - expected_active_ids = [ - int(message["_row_id"]) - for message in durable - if isinstance(message.get("_row_id"), int) - ] - scaffold, _ = split_user_originated_turn(durable_target) result = self._session_db.rewind_to_message( self.session_id, @@ -10262,7 +10259,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): # Walk backwards to the last *real* user message. Timeline bookkeeping # rows (display_kind set) are role=user but are not user turns — match - # CLI resume counting and list_recent_user_messages. Compaction + # CLI resume counting and user_originated_turn_view. Compaction # handoffs are excluded too (durable role=user, sometimes without # display_kind on legacy sessions; #80622). from agent.context_compressor import ( @@ -10363,7 +10360,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): # Walk backwards collecting the indices of the last N *real* user # messages (exclude display_kind timeline rows and compaction - # handoffs — same predicate as list_recent_user_messages, resume + # handoffs — same predicate as user_originated_turn_view, resume # turn counting, and /retry; #80622). from agent.context_compressor import ( history_before_user_originated_turn, diff --git a/gateway/session.py b/gateway/session.py index b283bfcce9..a2ea749919 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -4065,15 +4065,11 @@ class SessionStore: ) try: + expected_active_ids = self._db.get_active_message_ids(session_id) durable = self._db.get_messages_as_conversation( session_id, include_row_ids=True, ) - expected_active_ids = [ - int(message["_row_id"]) - for message in durable - if isinstance(message.get("_row_id"), int) - ] user_indices = [ index for index, message in enumerate(durable) @@ -4089,13 +4085,15 @@ class SessionStore: handoff, target_view = split_user_originated_turn(target) if target_view is None: return None - if require_retryable_composite: - if handoff is None: - return None - target_text = retryable_user_text(target_view.get("content")) + if require_retryable_composite and handoff is None: + return None except Exception as e: logger.debug("rewind_session: failed to resolve canonical target: %s", e) return None + if require_retryable_composite: + # Keep replay-policy failures distinct from persistence errors + # so /retry can explain why the selected carrier is unsafe. + target_text = retryable_user_text(target_view.get("content")) try: result = self._db.rewind_to_message( session_id, diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index f679724034..c5dbf582e1 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -2681,11 +2681,14 @@ class GatewaySlashCommandsMixin: # archive that row/tail and insert its pure scaffold atomically. # Plain turns keep the existing rewrite path below; #84078 owns # its separate archive_dropped/prefix-CAS semantics. - rewind_result = await self.async_session_store.rewind_session( - session_entry.session_id, - 1, - require_retryable_composite=True, - ) + try: + rewind_result = await self.async_session_store.rewind_session( + session_entry.session_id, + 1, + require_retryable_composite=True, + ) + except ValueError as exc: + return f"Cannot retry that message safely: {exc}" if rewind_result is None: return "Retry failed; transcript was not changed." # The store reselects and validates the latest carrier on the same diff --git a/hermes_cli/web_routers/sessions.py b/hermes_cli/web_routers/sessions.py index 2e88a3ea8e..933db5ed5d 100644 --- a/hermes_cli/web_routers/sessions.py +++ b/hermes_cli/web_routers/sessions.py @@ -642,23 +642,24 @@ async def get_session_messages( if result is None: raise HTTPException(status_code=404, detail="Session not found") sid, _limit, messages = result - from agent.context_compressor import split_user_originated_turn + from agent.compaction_display import project_compaction_message_for_display + from agent.context_compressor import is_compaction_summary_message projected_messages = [] for message in messages: - handoff, live_view = split_user_originated_turn(message) - if handoff is None: + if not is_compaction_summary_message(message): projected_messages.append(message) continue + display_view = project_compaction_message_for_display(message) projected = message.copy() - if live_view is None: + if display_view is None: if not projected.get("display_kind"): projected["display_kind"] = "hidden" else: # Keep the physical content for inspection/export compatibility; # Desktop consumes this display-only projection. A legacy hidden # wrapper must not hide a successfully recovered live ask. - projected["display_content"] = live_view.get("content") + projected["display_content"] = display_view.get("content") projected.pop("display_kind", None) projected_messages.append(projected) return { diff --git a/hermes_state.py b/hermes_state.py index 60c3984b29..b7c0304657 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -11745,6 +11745,22 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # Rewind (soft-delete) — see /rewind slash command + issue #21910 # ========================================================================= + def get_active_message_ids(self, session_id: str) -> List[int]: + """Return the ordered physical ids pinned by rewind CAS checks. + + Conversation projections intentionally omit legacy background-review + harness rows. Destructive rewinds must nevertheless pin every active + physical row so the caller snapshot matches the transaction-local + comparison in :meth:`rewind_to_message`. + """ + with self._read_ctx() as conn: + rows = conn.execute( + "SELECT id FROM messages " + "WHERE session_id = ? AND active = 1 ORDER BY id", + (session_id,), + ).fetchall() + return [int(row[0]) for row in rows] + @staticmethod def _active_transcript_counts(conn, session_id: str) -> tuple[int, int]: """Return active message/tool-call counts inside the caller's txn.""" diff --git a/tests/gateway/test_retry_replacement.py b/tests/gateway/test_retry_replacement.py index b3f2e8d96f..4dec6eba04 100644 --- a/tests/gateway/test_retry_replacement.py +++ b/tests/gateway/test_retry_replacement.py @@ -110,6 +110,32 @@ def test_rewind_session_keeps_pending_recovery_state_when_lease_rejects( assert session_id not in store._transcript_append_failures +def test_rewind_session_surfaces_unretryable_media_before_mutation( + tmp_path, monkeypatch +): + import hermes_state + + monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db") + store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig()) + session_id = "rewind-composite-media" + store._db.create_session(session_id=session_id, source="test") + store._db.append_message( + session_id, + "user", + [ + {"type": "text", "text": _composite_carrier()["content"]}, + {"type": "image_url", "image_url": {"url": "image"}}, + ], + ) + store._db.append_message(session_id, "assistant", "old answer") + before = store._db.get_messages(session_id, include_inactive=True) + + with pytest.raises(ValueError, match="media or unknown content"): + store.rewind_session(session_id, require_retryable_composite=True) + + assert store._db.get_messages(session_id, include_inactive=True) == before + + @pytest.mark.parametrize("operation", ["rewrite", "rewind"]) def test_transcript_mutation_serializes_pending_queue_drain( operation, tmp_path, monkeypatch @@ -383,6 +409,39 @@ async def test_gateway_retry_rejects_media_before_redispatch_or_token_reset(): facade.rewrite_transcript.assert_not_awaited() +@pytest.mark.asyncio +async def test_gateway_retry_preserves_composite_media_diagnostic_from_store(): + gw = GatewayRunner.__new__(GatewayRunner) + backing_store = MagicMock() + gw.session_store = backing_store + session_entry = SimpleNamespace(session_id="sid", last_prompt_tokens=123) + facade = SimpleNamespace( + _store=backing_store, + get_or_create_session=AsyncMock(return_value=session_entry), + load_transcript=AsyncMock( + return_value=[ + _composite_carrier(), + {"role": "assistant", "content": "old answer"}, + ] + ), + rewind_session=AsyncMock( + side_effect=ValueError("retry does not support media content") + ), + ) + gw._async_session_store = facade + gw._handle_message = AsyncMock() + + result = await gw._handle_retry_command( + MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock()) + ) + + assert result == ( + "Cannot retry that message safely: retry does not support media content" + ) + assert session_entry.last_prompt_tokens == 123 + gw._handle_message.assert_not_awaited() + + @pytest.mark.asyncio async def test_gateway_retry_stops_when_transcript_rewrite_fails(): gw = GatewayRunner.__new__(GatewayRunner) diff --git a/tests/gateway/test_undo_rewind_session.py b/tests/gateway/test_undo_rewind_session.py index 10df30a431..3f2f518c92 100644 --- a/tests/gateway/test_undo_rewind_session.py +++ b/tests/gateway/test_undo_rewind_session.py @@ -54,6 +54,35 @@ def test_rewind_n_turns(store): assert len(store.load_transcript(sid)) == 2 # q1,a1 +def test_rewind_pins_raw_active_ids_when_projection_hides_review_harness(store): + sid = _seed(store, "gw-review-harness", turns=2) + store._db.append_message( + sid, + "user", + "Review the conversation above and update the skill library safely", + ) + store._db.append_message(sid, "assistant", "curator-only reply") + + # Legacy background-review rows are intentionally absent from replay, but + # they remain physical active rows that the rewind CAS must pin. + assert [message["content"] for message in store.load_transcript(sid)] == [ + "q1", + "a1", + "q2", + "a2", + ] + + result = store.rewind_session(sid) + + assert result is not None + assert result["target_text"] == "q2" + assert result["rewound_count"] == 4 + assert [message["content"] for message in store.load_transcript(sid)] == [ + "q1", + "a1", + ] + + def test_rewind_fails_closed_when_transcript_changes_after_snapshot( store, monkeypatch ): @@ -83,3 +112,35 @@ def test_rewind_fails_closed_when_transcript_changes_after_snapshot( ] sibling.close() + +def test_rewind_fails_closed_when_new_turn_lands_after_id_snapshot( + store, monkeypatch +): + sid = _seed(store, "gw-snapshot-order", turns=2) + sibling = SessionDB(db_path=store._db.db_path) + original_load = store._db.get_messages_as_conversation + + def _load_then_append(*args, **kwargs): + snapshot = original_load(*args, **kwargs) + sibling.append_message(sid, "user", "q3-from-other-process") + sibling.append_message(sid, "assistant", "a3-from-other-process") + return snapshot + + monkeypatch.setattr(store._db, "get_messages_as_conversation", _load_then_append) + + assert store.rewind_session(sid) is None + + rows = store._db._conn.execute( + "SELECT content, active FROM messages " + "WHERE session_id = ? ORDER BY id", + (sid,), + ).fetchall() + assert [tuple(row) for row in rows] == [ + ("q1", 1), + ("a1", 1), + ("q2", 1), + ("a2", 1), + ("q3-from-other-process", 1), + ("a3-from-other-process", 1), + ] + sibling.close() diff --git a/tests/hermes_cli/test_web_server.py b/tests/hermes_cli/test_web_server.py index b9820c127c..0cfb7e927d 100644 --- a/tests/hermes_cli/test_web_server.py +++ b/tests/hermes_cli/test_web_server.py @@ -2048,6 +2048,8 @@ class TestWebServerEndpoints: from agent.context_compressor import ( HISTORICAL_TASK_HEADING, SUMMARY_PREFIX, + _MERGED_PRIOR_CONTEXT_HEADER, + _MERGED_SUMMARY_DELIMITER, _SUMMARY_END_MARKER, ) from hermes_state import SessionDB @@ -2057,6 +2059,11 @@ class TestWebServerEndpoints: f"{_SUMMARY_END_MARKER}" ) carrier = f"{handoff}\n\nREAL ASK" + assistant_carrier = ( + f"{_MERGED_PRIOR_CONTEXT_HEADER}\n" + "real completed answer\n\n" + f"{_MERGED_SUMMARY_DELIMITER}\n\n{handoff}" + ) db = SessionDB() try: db.create_session(session_id="compacted-carrier-display", source="desktop") @@ -2076,6 +2083,12 @@ class TestWebServerEndpoints: handoff, timestamp=124.0, ) + db.append_message( + "compacted-carrier-display", + "assistant", + assistant_carrier, + timestamp=125.0, + ) active_id = db.get_messages("compacted-carrier-display")[0]["id"] finally: db.close() @@ -2086,13 +2099,15 @@ class TestWebServerEndpoints: ) assert resp.status_code == 200 messages = resp.json()["messages"] - assert len(messages) == 2 + assert len(messages) == 3 assert messages[0]["id"] == active_id assert messages[0]["content"] == carrier assert messages[0]["display_content"] == "REAL ASK" assert not messages[0].get("display_kind") assert messages[1]["content"] == handoff assert messages[1]["display_kind"] == "hidden" + assert messages[2]["content"] == assistant_carrier + assert messages[2]["display_content"] == "real completed answer" def test_get_session_messages_latest_page_with_compacted_rows(self): """The desktop's real read path (getLatestSessionMessages: limit + diff --git a/tests/tui_gateway/test_composite_carrier_rewind.py b/tests/tui_gateway/test_composite_carrier_rewind.py index bef256746d..394d992b69 100644 --- a/tests/tui_gateway/test_composite_carrier_rewind.py +++ b/tests/tui_gateway/test_composite_carrier_rewind.py @@ -403,6 +403,36 @@ def test_retry_preserves_literal_media_like_text(carrier_session): _assert_scaffold_preserved(db, session_key, session) +def test_retry_rejects_durable_media_before_rewind_when_warm_view_is_text( + carrier_session, +): + db, install = carrier_session + carrier = _composite_carrier() + handoff = carrier["content"].rsplit("\n\nREAL ASK", 1)[0] + durable_carrier = carrier.copy() + durable_carrier["content"] = [ + {"type": "text", "text": handoff}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,x"}}, + ] + sid, session_key, session = install( + [durable_carrier, {"role": "assistant", "content": "failed"}] + ) + # The warm projection can be a degraded text-only view that compares equal + # to the durable media payload. Durable retryability must still be checked + # before the physical carrier and tail are archived. + warm_carrier = carrier.copy() + warm_carrier["content"] = handoff + "\n\n[screenshot]" + session["history"][0] = warm_carrier + session["agent"]._session_messages = list(session["history"]) + before_history = [message.copy() for message in session["history"]] + + response = _dispatch(sid, "retry") + + assert response["error"]["code"] == 4018 + assert session["history"] == before_history + assert len(db.get_messages_as_conversation(session_key)) == 2 + + def test_retry_rejects_pending_attachments_before_mutating_history(carrier_session): db, install = carrier_session sid, session_key, session = install( diff --git a/tui_gateway/methods_session.py b/tui_gateway/methods_session.py index 82637f4dc1..229f474300 100644 --- a/tui_gateway/methods_session.py +++ b/tui_gateway/methods_session.py @@ -2694,7 +2694,7 @@ def _(rid, params: dict) -> dict: # (async_delegation_complete, model_switch, …) or compaction # handoffs as the undo target — so session.undo removed # bookkeeping instead of the last exchange (#80622). - # Match list_recent_user_messages / CLI turn counting. + # Match user_originated_turn_view / CLI turn counting. from agent.context_compressor import user_originated_turn_view user_indices = [ diff --git a/tui_gateway/methods_tools.py b/tui_gateway/methods_tools.py index 449b336af0..f153bc98a0 100644 --- a/tui_gateway/methods_tools.py +++ b/tui_gateway/methods_tools.py @@ -744,8 +744,14 @@ def _(rid, params: dict) -> dict: return _err(rid, 4018, str(exc)) try: _active, durable_live_view, _rewound_count = ( - _rewind_active_session_history(session, len(user_indices) - 1) + _rewind_active_session_history( + session, + len(user_indices) - 1, + require_retryable=True, + ) ) + except ValueError as exc: + return _err(rid, 4018, str(exc)) except Exception as exc: return _err(rid, 5008, f"retry: failed to persist history: {exc}") content = retryable_user_text(durable_live_view.get("content")) diff --git a/tui_gateway/server.py b/tui_gateway/server.py index 38a53b2a8a..f767e60e36 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -3155,7 +3155,10 @@ def _session_db(session: dict): def _rewind_active_session_history( - session: dict, user_ordinal: int + session: dict, + user_ordinal: int, + *, + require_retryable: bool = False, ) -> tuple[list[dict], dict, int]: """Rewind one canonical user turn while retaining carrier scaffolding. @@ -3167,6 +3170,7 @@ def _rewind_active_session_history( """ from agent.context_compressor import ( history_before_user_originated_turn, + retryable_user_text, split_user_originated_turn, user_originated_turn_view, ) @@ -3213,6 +3217,7 @@ def _rewind_active_session_history( with _session_db(session) as db: if db is None: raise RuntimeError("session database is unavailable") + expected_active_ids = db.get_active_message_ids(session_key) durable = db.get_messages_as_conversation( session_key, include_row_ids=True, @@ -3238,11 +3243,8 @@ def _rewind_active_session_history( target_row_id = durable_target.get("_row_id") if not isinstance(target_row_id, int): raise RuntimeError("rewind target has no durable row identity") - expected_active_ids = [ - int(message["_row_id"]) - for message in durable - if isinstance(message.get("_row_id"), int) - ] + if require_retryable: + retryable_user_text(durable_live_view.get("content")) scaffold, _ = split_user_originated_turn(durable_target) result = db.rewind_to_message( session_key, @@ -3278,6 +3280,8 @@ def _rewind_active_session_history( live_view = durable_live_view rewound_count = int(result.get("rewound_count", 0)) persisted = True + elif require_retryable: + retryable_user_text(live_view.get("content")) installed = [message.copy() for message in installed] session["history"] = installed @@ -7854,20 +7858,6 @@ def _history_to_messages(history: list[dict]) -> list[dict]: role = m.get("role") if role not in {"user", "assistant", "tool", "system"}: continue - if role == "user": - from agent.context_compressor import ( - is_compaction_summary_message, - user_originated_turn_view, - ) - - if is_compaction_summary_message(m): - carrier_row_id = m.get("_row_id") - live_view = user_originated_turn_view(m) - if live_view is None: - continue - if carrier_row_id is not None: - live_view["_row_id"] = carrier_row_id - m = live_view # An explicit display_kind="hidden" row is model-facing scaffolding # (compaction references, interrupted-turn checkpoints). The string # sniff below only catches the "[System:" convention; honor the