diff --git a/hermes_state_rewind.py b/hermes_state_rewind.py index 1ea08214d1..37ce926a6f 100644 --- a/hermes_state_rewind.py +++ b/hermes_state_rewind.py @@ -21,7 +21,7 @@ class RewindTargetUnavailableError(ValueError): class RewindOutcome: prefix: List[Dict[str, Any]] # history to install: the warm prefix when ``warm_history`` was given, else durable live_view: Dict[str, Any] # canonical live projection of the rewound turn (prefill / retry source) - live_text: str + live_text: str # lossless retry text when ``require_retryable``, else the display flattening (prefill) rewound_count: int turns_undone: int @@ -85,8 +85,9 @@ class SessionRewindMixin: prefix, warm_live_view = history_before_user_originated_turn(warm, warm_user[user_ordinal]) if _comparison_content(live_view) != _comparison_content(warm_live_view): raise RuntimeError(_HISTORY_CHANGED) - if require_retryable: - retryable_user_text(live_view.get("content")) + # Retry re-sends the stored bytes: ``"".join`` of the text parts, never the "\n"-joined display + # flattening (wire bytes == stored bytes; ``"ab"`` must not come back as ``"a\nb"``). + live_text = retryable_user_text(live_view.get("content")) if require_retryable else None target_row_id = target.get("_row_id") if not isinstance(target_row_id, int): raise RuntimeError("rewind target has no durable row identity") @@ -113,5 +114,6 @@ class SessionRewindMixin: if isinstance(row_id := durable_message.get("_row_id"), int): warm["_row_id"] = row_id return RewindOutcome( - prefix=prefix, live_view=live_view, live_text=flatten_message_text(live_view.get("content")), + prefix=prefix, live_view=live_view, + live_text=live_text if live_text is not None else flatten_message_text(live_view.get("content")), rewound_count=int(result.get("rewound_count", 0)), turns_undone=len(durable_user) - user_ordinal) diff --git a/tests/gateway/test_retry_replacement.py b/tests/gateway/test_retry_replacement.py index ed3f35f32f..ca5bcf1050 100644 --- a/tests/gateway/test_retry_replacement.py +++ b/tests/gateway/test_retry_replacement.py @@ -110,6 +110,37 @@ def test_rewind_session_keeps_pending_recovery_state_when_lease_rejects( assert session_id not in store._transcript_append_failures +def test_rewind_session_retry_text_is_the_stored_bytes_for_multipart_carriers( + tmp_path, monkeypatch +): + """/retry re-sends exactly what was stored: a carrier split across text parts comes back as the + parts' concatenation (``"".join``), not the "\\n"-joined display flattening.""" + import hermes_state + from agent.context_compressor import retryable_user_text + + monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db") + store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig()) + session_id = "rewind-composite-multipart" + store._db.create_session(session_id=session_id, source="test") + parts = [ + {"type": "text", "text": _composite_carrier("REAL")["content"]}, + {"type": "text", "text": " ASK"}, + {"type": "text", "text": "\ncontinued"}, + ] + store._db.append_message(session_id, "user", parts) + store._db.append_message(session_id, "assistant", "old answer") + stored = store._db.get_messages_as_conversation(session_id)[0]["content"] + assert isinstance(stored, list) and len(stored) == 3 + + result = store.rewind_session(session_id, require_retryable_composite=True) + + assert result is not None + expected = "".join(p["text"] for p in stored).split(_SUMMARY_END_MARKER, 1)[1].lstrip("\n") + assert result["target_text"] == "REAL ASK\ncontinued" == expected + assert result["target_text"] == retryable_user_text( + [{"type": "text", "text": "REAL"}, {"type": "text", "text": " ASK"}, {"type": "text", "text": "\ncontinued"}]) + + def test_rewind_session_surfaces_unretryable_media_before_mutation( tmp_path, monkeypatch ):