From c3de99bbf05994dcb6a1f334430b191f67e3c0af Mon Sep 17 00:00:00 2001 From: teknium1 <127238744+teknium1@users.noreply.github.com> Date: Sat, 12 Sep 2026 21:10:03 -0700 Subject: [PATCH] refactor(state): one carrier-aware user-turn rewind behind CLI /undo, /retry, gateway and TUI CLI `_rewind_persisted_user_turn`, TUI `_rewind_active_session_history` and gateway `rewind_session` each re-ran get_active_message_ids -> get_messages_as_conversation -> split_user_originated_turn -> rewind_to_message with their own warm/durable comparison helpers and three different out-of-range contracts (RuntimeError / ValueError / None). The durable transcript is the authority for a rewind, so the implementation now lives with the data: `SessionDB.rewind_user_turn` (hermes_state_rewind.py) with one typed out-of-range error (`RewindTargetUnavailableError`). Surfaces keep only lock, eviction and rendering glue and map that error to their own message. --- docs/session-lifecycle.md | 2 +- gateway/session_transcript.py | 72 ++---- hermes_cli/cli_session_mixin.py | 215 +++--------------- hermes_state.py | 3 +- hermes_state_rewind.py | 117 ++++++++++ tests/cli/test_cli_retry.py | 5 +- .../test_rewind_surfaces_invariant.py | 98 ++++++++ tui_gateway/session_workdir.py | 86 ++----- 8 files changed, 284 insertions(+), 314 deletions(-) create mode 100644 hermes_state_rewind.py create mode 100644 tests/hermes_state/test_rewind_surfaces_invariant.py diff --git a/docs/session-lifecycle.md b/docs/session-lifecycle.md index ddce596be8..7e22ad6be6 100644 --- a/docs/session-lifecycle.md +++ b/docs/session-lifecycle.md @@ -177,7 +177,7 @@ SessionStore(sessions_dir: Path, config: GatewayConfig, has_active_processes_fn= | `append_to_transcript(session_id, message, skip_db=False)` | Append a message to SQLite transcript. `skip_db=True` prevents duplicate writes when the agent already persisted. | | `rewrite_transcript(session_id, messages)` | Full replacement of session transcript (used by `/retry`, `/undo`, `/compress`). | | `load_transcript(session_id)` | Load all messages from a session's SQLite transcript. | -| `rewind_session(session_id, n=1)` | Back up `n` user turns via soft-delete (keeps audit trail). Returns `{rewound_count, turns_undone, target_text}`. | +| `rewind_session(session_id, n=1)` | Back up `n` user turns via soft-delete (keeps audit trail); thin wrapper over `SessionDB.rewind_user_turn` (`hermes_state_rewind.py`), the one rewind shared with CLI `/undo`/`/retry` and the TUI. Returns `{rewound_count, turns_undone, target_text}`. | ### Internal Helpers diff --git a/gateway/session_transcript.py b/gateway/session_transcript.py index ba4153d4e0..e24fe43a48 100644 --- a/gateway/session_transcript.py +++ b/gateway/session_transcript.py @@ -25,14 +25,6 @@ class TranscriptReadError(RuntimeError): super().__init__(f"transcript read failed for session {session_id}") -def _plain_text(content) -> str: - """Text of a message content (str or text-part list); "" for anything else.""" - if isinstance(content, list): - parts = [p.get("text", "") for p in content if isinstance(p, dict) and p.get("type") == "text"] - return "\n".join(t for t in parts if t) - return content if isinstance(content, str) else "" - - def _spool_dropped(session_id: str, message: Dict[str, Any]): """Spool one evicted/undeliverable message to disk (same machinery as the shutdown flush, so it is replayed after DB recovery); path or None.""" @@ -516,60 +508,28 @@ class SessionTranscriptMixin: self, session_id: str, n: int = 1, *, require_retryable_composite: bool = False, ) -> Optional[Dict[str, Any]]: """Back up ``n`` user turns via soft-delete (``active=0``), mirroring CLI ``/undo [N]``. - Returns ``{"rewound_count", "turns_undone", "target_text"}`` or ``None`` (no DB / no user - turn); ``n`` clamps to the oldest user turn. ``require_retryable_composite`` is the gateway - ``/retry`` guard: the selected turn must be a composite carrier whose live payload is - losslessly replayable as text before anything changes.""" + Returns ``{"rewound_count", "turns_undone", "target_text"}`` or ``None`` (no DB / no rewindable + turn / persistence failure); ``n`` clamps to the oldest user turn. ``require_retryable_composite`` + is the gateway ``/retry`` guard: the selected turn must be a composite carrier whose live payload + is losslessly replayable as text — that replay-policy ``ValueError`` propagates so /retry can + explain why the carrier is unsafe.""" db = self._db_for_session_id(session_id) if not db: return None + from hermes_state_rewind import RewindTargetUnavailableError with self._get_transcript_drain_lock(): - n = max(n, 1) - from agent.context_compressor import ( - retryable_user_text, split_user_originated_turn, user_originated_turn_view, - ) try: - expected_active_ids = db.get_active_message_ids(session_id) - durable = db.get_messages_as_conversation(session_id, include_row_ids=True) - user_indices = [ - index for index, message in enumerate(durable) - if user_originated_turn_view(message) is not None - ] - if not user_indices: - return None - turns_undone = min(n, len(user_indices)) - target = durable[user_indices[-turns_undone]] - target_id = target.get("_row_id") - if not isinstance(target_id, int): - return None - handoff, target_view = split_user_originated_turn(target) - if target_view is None: - return None - 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) + outcome = db.rewind_user_turn( + session_id, -max(n, 1), require_retryable=require_retryable_composite, + require_composite=require_retryable_composite) + except RewindTargetUnavailableError as e: + logger.debug("rewind_session: %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 = db.rewind_to_message( - session_id, target_id, preserve_compaction_handoff=handoff is not None, - expected_active_ids=expected_active_ids, - expected_target_content=target_view.get("content")) + except ValueError: + raise except Exception as e: - prefix = "" if isinstance(e, ValueError) else "rewind_to_message failed: " - logger.debug("rewind_session: %s%s", prefix, e) + logger.debug("rewind_session: rewind failed: %s", e) return None self._clear_dirty_transcript(session_id) - # ``target_view`` is the live projection; a composite carrier's raw row holds the - # summary wrapper and must not be echoed as prompt. - if not require_retryable_composite: - target_text = _plain_text(target_view.get("content") or "") - return { - "rewound_count": result.get("rewound_count", 0), - "turns_undone": turns_undone, - "target_text": target_text, - } + return {"rewound_count": outcome.rewound_count, "turns_undone": outcome.turns_undone, + "target_text": outcome.live_text} diff --git a/hermes_cli/cli_session_mixin.py b/hermes_cli/cli_session_mixin.py index b9ae6400e5..b807483cf1 100644 --- a/hermes_cli/cli_session_mixin.py +++ b/hermes_cli/cli_session_mixin.py @@ -709,87 +709,6 @@ class CLISessionMixin: except Exception as e: print(f"(x_x) Failed to save: {e}") - def _rewind_persisted_user_turn( - self, - *, - warm_history: List[Dict[str, Any]], - user_ordinal: int, - warm_live_view: Dict[str, Any], - ) -> tuple[List[Dict[str, Any]], Dict[str, Any], Dict[str, Any]]: - """Bind one warm user ordinal to a durable row and rewind it atomically.""" - if self._session_db is None or not self.session_id: - raise RuntimeError("session database is unavailable") - - from agent.context_compressor import ( - history_before_user_originated_turn, - split_user_originated_turn, - user_originated_turn_view) - from agent.memory_manager import sanitize_context - from agent.tool_dispatch_helpers import _is_multimodal_tool_result, _multimodal_text_summary - from agent.session_persistence import _is_ephemeral_scaffolding - - def _persistence_content(content: Any) -> Any: - """Project warm content exactly as the session DB flush does.""" - if _is_multimodal_tool_result(content): - return _multimodal_text_summary(content) - if isinstance(content, list): - text_parts = [] - for part in content: - if isinstance(part, dict) and part.get("type") == "text": - text_parts.append(str(part.get("text", ""))) - elif isinstance(part, dict) and part.get("type") in { - "image", "image_url", "input_image"}: - text_parts.append("[screenshot]") - return "\n".join(text_parts) if text_parts else None - return content - - def _comparison_content(message: Dict[str, Any]) -> Any: - content = _persistence_content(message.get("content")) - if message.get("role") in {"user", "assistant"} and isinstance(content, str): - return sanitize_context(content).strip() - return content - - def _user_indices(messages): - return [i for i, m in enumerate(messages) if user_originated_turn_view(m) is not None] - - changed = RuntimeError("session history changed before the rewind could be persisted") - 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) - warm_persistence_history = [m for m in warm_history if not _is_ephemeral_scaffolding(m)] - warm_user_indices = _user_indices(warm_persistence_history) - durable_user_indices = _user_indices(durable) - if len(durable_user_indices) != len(warm_user_indices): - raise changed - if user_ordinal < 0 or user_ordinal >= len(durable_user_indices): - raise RuntimeError("persisted rewind target is no longer available") - - warm_prefix, _ = history_before_user_originated_turn( - warm_persistence_history, warm_user_indices[user_ordinal]) - durable_target_index = durable_user_indices[user_ordinal] - durable_target = durable[durable_target_index] - durable_prefix, durable_live_view = history_before_user_originated_turn( - durable, durable_target_index) - if _comparison_content(durable_live_view) != _comparison_content(warm_live_view): - raise changed - target_row_id = durable_target.get("_row_id") - if not isinstance(target_row_id, int): - raise RuntimeError("persisted rewind target has no row identity") - scaffold, _ = split_user_originated_turn(durable_target) - result = self._session_db.rewind_to_message( - self.session_id, target_row_id, - preserve_compaction_handoff=scaffold is not None, - expected_active_ids=expected_active_ids, - expected_target_content=durable_live_view.get("content")) - if scaffold is not None: - replacement_id = result.get("replacement_message_id") - if not isinstance(replacement_id, int) or not durable_prefix: - raise RuntimeError("rewind did not retain its compaction handoff") - durable_prefix[-1]["_row_id"] = replacement_id - durable_prefix[-1]["_db_persisted"] = True - warm_prefix[-1] = durable_prefix[-1] - return warm_prefix, durable_live_view, result - def _publish_truncated_history(self, truncated: list, *, invalidate_prompt: bool) -> None: """Install a rewound history and mirror it onto the agent (flush index reset so the next turn re-flushes from the truncated head).""" @@ -844,10 +763,8 @@ class CLISessionMixin: # physical carrier split so archived original + retained scaffold commit atomically. if self._session_db is not None and self.session_id: try: - truncated, _, _ = self._rewind_persisted_user_turn( - warm_history=warm_history, - user_ordinal=len(user_indices) - 1, - warm_live_view=live_view) + truncated = self._session_db.rewind_user_turn( + self.session_id, -1, warm_history=warm_history, require_retryable=True).prefix except Exception as exc: print(f"(x_x) Retry rewind failed; history was not changed: {exc}") return None @@ -889,15 +806,12 @@ class CLISessionMixin: rewound_rows = 0 if self._session_db is not None and self.session_id: try: - truncated, durable_live_view, result = self._rewind_persisted_user_turn( - warm_history=warm_history, - user_ordinal=target_ordinal, - warm_live_view=live_view) + outcome = self._session_db.rewind_user_turn( + self.session_id, target_ordinal, warm_history=warm_history) + truncated = outcome.prefix # Canonical editable prefill: the raw carrier holds the reference-summary wrapper. - durable_text = self._undo_content_to_text(durable_live_view.get("content")) - if durable_text: - removed_text = durable_text - rewound_rows = result.get("rewound_count", 0) + removed_text = outcome.live_text or removed_text + rewound_rows = outcome.rewound_count except Exception as e: logger.debug("undo: durable rewind failed: %s", e) print(f"(x_x) Undo failed; history was not changed: {e}") @@ -1034,103 +948,47 @@ class CLISessionMixin: the summariser what to preserve while discarding the rest more aggressively. * ``/compress here [N]`` — boundary-aware: summarize everything except the most recent ``N`` exchanges (default 2), kept verbatim. + * ``--preview`` reports what would happen and changes nothing. No ``compression_enabled`` gate: that flag disables *automatic* compaction only, and the context-overflow error path directs users here when it is off. """ - if len(self.conversation_history or ()) < 4: - print("(._.) Not enough conversation to compress (need at least 4 messages).") + from agent.conversation_compression import finalize_context_engine_compression_notification + from agent.conversation_compression_manual import ( + AGGRESSIVE_UNSUPPORTED, MIN_MESSAGES, compress_now, parse_compress_args, render_compress_result) + + if len(self.conversation_history or ()) < MIN_MESSAGES: + print(f"(._.) Not enough conversation to compress (need at least {MIN_MESSAGES} messages).") return if not self.agent: print("(._.) No active agent -- send a message first.") return - - from hermes_cli.partial_compress import ( - extract_compress_flags, parse_partial_compress_args, rejoin_compressed_head_and_tail, - split_history_for_partial_compress, summarize_compress_preview) - from agent.conversation_compression import finalize_context_engine_compression_notification - from agent.model_metadata import estimate_request_tokens_rough - _parts = (cmd_original or "").strip().split(None, 1) - raw_args = _parts[1].strip() if len(_parts) > 1 else "" - # Strip --preview/--dry-run/--aggressive before positional parsing. - raw_args, preview, aggressive = extract_compress_flags(raw_args) - partial, keep_last, focus_topic = parse_partial_compress_args(raw_args) - focus_topic = focus_topic or "" - - if aggressive: - # LLM-free hard truncation would need its own persistence path outside the - # guarded _compress_context rotation; surface that instead of mis-parsing. - print("(._.) --aggressive is not supported; use '/compress here [N]' " - "to keep only recent exchanges, or /undo to drop turns.") - if not preview: + request = parse_compress_args(_parts[1] if len(_parts) > 1 else "") + if request.aggressive: + print(f"(._.) {AGGRESSIVE_UNSUPPORTED}") + if not request.preview: return - - # Include system prompt + tool schemas in estimates — a transcript-only number - # understates real request pressure and can even appear to grow after compression. - _estimate_kw = { - "system_prompt": getattr(self.agent, "_cached_system_prompt", "") or "", - "tools": getattr(self.agent, "tools", None) or None} - if preview: - approx_tokens = estimate_request_tokens_rough(self.conversation_history, **_estimate_kw) - report = summarize_compress_preview( - self.conversation_history, partial, keep_last, focus_topic or None, approx_tokens) - for line in report["lines"]: + if request.preview: + for line in render_compress_result(compress_now(self.agent, self.conversation_history, request)): print(f"🗜️ {line}") return original_count = len(self.conversation_history) with self._busy_command("Compressing context...", blocks_input=False): try: - from agent.manual_compression_feedback import summarize_manual_compression - original_history = list(self.conversation_history) - - # Boundary-aware split: only the head is summarized. A degenerate split - # (nothing to keep / no head) falls back to full compression. - tail: list = [] - head = original_history - if partial: - head, tail = split_history_for_partial_compress(original_history, keep_last) - if not tail: - partial = False - head = original_history - - approx_tokens = estimate_request_tokens_rough(original_history, **_estimate_kw) - if partial: - print(f"🗜️ Summarizing up to here: compressing {len(head)} of " - f"{original_count} messages (~{approx_tokens:,} tokens), " - f"keeping last {keep_last} exchange(s) verbatim...") - elif focus_topic: - print(f"🗜️ Compressing {original_count} messages (~{approx_tokens:,} tokens), " - f"focus: \"{focus_topic}\"...") + if request.partial: + print(f"🗜️ Summarizing up to here: {original_count} messages, " + f"keeping last {request.keep_last} exchange(s) verbatim...") + elif request.focus_topic: + print(f"🗜️ Compressing {original_count} messages, focus: \"{request.focus_topic}\"...") else: - print(f"🗜️ Compressing {original_count} messages (~{approx_tokens:,} tokens)...") - - # system_message=None so _compress_context rebuilds the prompt from scratch; - # passing _cached_system_prompt duplicated the identity block. - # Passing _cached_system_prompt caused duplication because _build_system_prompt appends - # system_message to prompt_parts which already contain the agent identity — resulting in the - # identity block appearing twice (issue #15281). - compressed, _ = self.agent._compress_context( - head, None, approx_tokens=approx_tokens, focus_topic=focus_topic or None, - force=True, defer_context_engine_notification=True) - - # Unchanged because a concurrent compression lock is held: say so instead of - # the misleading "No changes" no-op text. Type-pinned check (is True / str) — - # a bare truthiness test is fooled by MagicMock auto-attributes on test doubles. - _lock_skip_signal = getattr(self.agent, "_compression_skipped_due_to_lock", None) - if _lock_skip_signal is True or isinstance(_lock_skip_signal, str): - from agent.manual_compression_feedback import describe_compression_lock_skip - print( - " " + describe_compression_lock_skip(self.agent._compression_skipped_due_to_lock) - ) - self.agent._compression_skipped_due_to_lock = None - # No boundary committed → discard the deferred notification (exactly-once). - finalize_context_engine_compression_notification(self.agent, committed=False) + print(f"🗜️ Compressing {original_count} messages...") + result = compress_now(self.agent, self.conversation_history, request) + if result.status != "compressed": + for line in render_compress_result(result): + print(f" {line}") return - - if partial and tail: - compressed = rejoin_compressed_head_and_tail(compressed, tail) - self.conversation_history = compressed + self.conversation_history = result.after_messages # _compress_context ends the old session and creates a child session on the # agent. Sync the CLI's session_id so /status, /resume, exit summary and title # generation point at the live continuation, not the ended parent. @@ -1142,15 +1000,8 @@ class CLISessionMixin: # Persist the new handoff from offset 0 so resume can recover it after exit. self.agent._flush_messages_to_session_db(self.conversation_history, None) finalize_context_engine_compression_notification(self.agent, committed=True) - new_tokens = estimate_request_tokens_rough( - self.conversation_history, **_estimate_kw) - summary = summarize_manual_compression( - original_history, self.conversation_history, approx_tokens, new_tokens, - compression_state=getattr(self.agent, "context_compressor", None)) - if ( - summary.get("aborted") - or summary.get("fallback_used") - or summary.get("refused_would_grow")): + summary = result.summary + if summary.get("aborted") or summary.get("fallback_used") or summary.get("refused_would_grow"): icon = "⚠️" else: icon = "🗜️" if summary["noop"] else "✅" diff --git a/hermes_state.py b/hermes_state.py index c69ea1e1a1..5d5cdcc6c6 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -53,6 +53,7 @@ from hermes_state_dbfile import ( RetiredGenerationCaptureError, capture_retired_wal_generation, refuse_deleted_wal_generation, ) from hermes_state_messages import SessionMessagesMixin +from hermes_state_rewind import SessionRewindMixin from hermes_state_wal import ( _WAL_INCOMPAT_MARKERS, _on_disk_journal_mode, apply_database_pragmas, apply_wal_with_fallback, ) @@ -392,7 +393,7 @@ class SessionDB( SessionSessionsMixin, SessionFtsSetupMixin, SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin, SessionTelegramTopicsMixin, SessionCompressionMixin, SessionGatewayMixin, SessionMaintenanceMixin, SessionUsageMixin, SessionTitlesMixin, - SessionMessagesMixin, + SessionMessagesMixin, SessionRewindMixin, ): """SQLite-backed session storage with FTS5 search; many reader threads, one writer (WAL).""" diff --git a/hermes_state_rewind.py b/hermes_state_rewind.py new file mode 100644 index 0000000000..1ea08214d1 --- /dev/null +++ b/hermes_state_rewind.py @@ -0,0 +1,117 @@ +"""Carrier-aware user-turn rewind (``/undo``, ``/retry``) — the ONE implementation behind the CLI, the +gateway and the TUI. Rewind is a persisted-history operation: the durable transcript is the authority, +the warm (in-memory) history only has to agree with it. A composite compaction carrier (retained +summary + live human ask in one row) keeps its hidden handoff scaffold as the new head.""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Dict, List, Optional + +_HISTORY_CHANGED = "session history changed before the rewind could be persisted" + + +class RewindTargetUnavailableError(ValueError): + """The requested user turn is not a rewindable target of the active transcript: no user turns, an + ordinal past the newest one, a row that is not user-originated, or a plain turn where the caller + required a compaction carrier. Surfaces map this to their own "nothing to undo" message.""" + + +@dataclass +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 + rewound_count: int + turns_undone: int + + +def _user_indices(messages: List[Dict[str, Any]]) -> List[int]: + from agent.context_compressor import user_originated_turn_view + return [i for i, m in enumerate(messages) if user_originated_turn_view(m) is not None] + + +def _comparison_content(message: Dict[str, Any]) -> Any: + """Project content the way the durable row stores it (flush projection, then the read-side sanitize) so a + warm row and its durable twin compare equal.""" + from agent.memory_manager import sanitize_context + from agent.session_persistence import _durable_content + content = _durable_content(message.get("content")) + if message.get("role") in {"user", "assistant"} and isinstance(content, str): + return sanitize_context(content).strip() + return content + + +class SessionRewindMixin: + """``SessionDB`` mixin: soft-delete from one user turn onward, validated against the warm history.""" + + def rewind_user_turn( + self, session_id: str, user_ordinal: int, *, warm_history: Optional[List[Dict[str, Any]]] = None, + require_retryable: bool = False, require_composite: bool = False, + ) -> RewindOutcome: + """Rewind the active transcript to just before user turn ``user_ordinal`` (0 = oldest; negative counts + back from the newest and clamps to the oldest, so ``-n`` is ``/undo n``). ``warm_history`` (CLI/TUI): + the in-memory view must have the same user turns and the same live target text as the durable + transcript, else ``RuntimeError`` and nothing changes; its (richer) prefix is what gets installed. + ``require_retryable``: the live payload must be losslessly replayable as text (``ValueError`` from + :func:`retryable_user_text` before any write). ``require_composite``: the target must be a compaction + carrier. Out-of-range / wrong-shape targets raise :class:`RewindTargetUnavailableError`.""" + from agent.context_compressor import ( + _DB_PERSISTED_MARKER, history_before_user_originated_turn, retryable_user_text, + split_user_originated_turn) + from agent.message_content import flatten_message_text + from agent.session_persistence import _is_ephemeral_scaffolding + + expected_active_ids = self.get_active_message_ids(session_id) + durable = self.get_messages_as_conversation(session_id, include_row_ids=True) + durable_user = _user_indices(durable) + if user_ordinal < 0: + user_ordinal = max(len(durable_user) + user_ordinal, 0) + if user_ordinal >= len(durable_user): + raise RewindTargetUnavailableError("target user message is no longer in session history") + target_index = durable_user[user_ordinal] + target = durable[target_index] + durable_prefix, live_view = history_before_user_originated_turn(durable, target_index) + scaffold, _ = split_user_originated_turn(target) + if require_composite and scaffold is None: + raise RewindTargetUnavailableError("target user message is not a compaction carrier") + + prefix = durable_prefix + if warm_history is not None: + warm = [m for m in warm_history if not _is_ephemeral_scaffolding(m)] + warm_user = _user_indices(warm) + if len(warm_user) != len(durable_user): + raise RuntimeError(_HISTORY_CHANGED) + 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")) + target_row_id = target.get("_row_id") + if not isinstance(target_row_id, int): + raise RuntimeError("rewind target has no durable row identity") + try: + result = self.rewind_to_message( + session_id, target_row_id, preserve_compaction_handoff=scaffold is not None, + expected_active_ids=expected_active_ids, expected_target_content=live_view.get("content")) + except ValueError as exc: # target vanished / changed role under us: same class of failure as out-of-range + raise RewindTargetUnavailableError(str(exc)) from exc + if scaffold is not None: + replacement_id = result.get("replacement_message_id") + if not isinstance(replacement_id, int) or not durable_prefix: + raise RuntimeError("rewind did not retain its compaction handoff") + durable_prefix[-1].update({"_row_id": replacement_id, _DB_PERSISTED_MARKER: True}) + prefix[-1] = durable_prefix[-1] + if prefix is not durable_prefix and len(prefix) == len(durable_prefix) and all( + warm.get("role") == durable_message.get("role") + and bool(warm.get("display_kind")) == bool(durable_message.get("display_kind")) + and _comparison_content(warm) == _comparison_content(durable_message) + for warm, durable_message in zip(prefix, durable_prefix) + ): + # Clients address follow-ups by durable row id: keep the richer warm content, adopt the identities. + for warm, durable_message in zip(prefix, durable_prefix): + 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")), + rewound_count=int(result.get("rewound_count", 0)), turns_undone=len(durable_user) - user_ordinal) diff --git a/tests/cli/test_cli_retry.py b/tests/cli/test_cli_retry.py index 8386184345..a4615f1f51 100644 --- a/tests/cli/test_cli_retry.py +++ b/tests/cli/test_cli_retry.py @@ -319,14 +319,13 @@ def test_retry_last_rejects_media_before_db_or_memory_mutation(): assert cli.retry_last() is None assert cli.conversation_history is history - db.get_messages_as_conversation.assert_not_called() - db.rewind_to_message.assert_not_called() + db.rewind_user_turn.assert_not_called() def test_retry_last_db_failure_leaves_warm_history_unchanged(): cli = _make_cli() db = MagicMock() - db.get_messages_as_conversation.side_effect = OSError("db unavailable") + db.rewind_user_turn.side_effect = OSError("db unavailable") cli._session_db = db history = [ {"role": "user", "content": "retry me"}, diff --git a/tests/hermes_state/test_rewind_surfaces_invariant.py b/tests/hermes_state/test_rewind_surfaces_invariant.py new file mode 100644 index 0000000000..0a323a746a --- /dev/null +++ b/tests/hermes_state/test_rewind_surfaces_invariant.py @@ -0,0 +1,98 @@ +"""Invariant: every /undo + /retry surface persists through ``SessionDB.rewind_user_turn``, so the same +session rewound through the CLI, the gateway store, or the TUI helper leaves the IDENTICAL active +message set, and an out-of-range target changes nothing on every surface.""" + +from __future__ import annotations + +import threading +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from gateway.config import GatewayConfig +from gateway.session import SessionStore +from hermes_state import SessionDB +from hermes_state_rewind import RewindTargetUnavailableError + +SURFACES = ("cli", "gateway", "tui") + + +def _seed(db: SessionDB, sid: str, turns: int = 3) -> None: + db.create_session(sid, source="cli") + for i in range(1, turns + 1): + db.append_message(sid, "user", f"q{i}") + db.append_message(sid, "assistant", f"a{i}") + + +def _active_rows(db: SessionDB, sid: str): + return [tuple(r) for r in db._conn.execute( + "SELECT id, role, content, active FROM messages WHERE session_id = ? ORDER BY id", (sid,)).fetchall()] + + +def _rewind_via(surface: str, db: SessionDB, sid: str, n: int): + """Back up ``n`` user turns through one surface; returns a truthy outcome or raises/None on refusal.""" + if surface == "gateway": + store = SessionStore.__new__(SessionStore) + store._db_for_session_id = lambda _sid: db + store._lazy = lambda name, factory: factory() + store._clear_dirty_transcript = lambda _sid: None + return store.rewind_session(sid, n) + warm = db.get_messages_as_conversation(sid) + user_turns = sum(1 for m in warm if m.get("role") == "user") + ordinal = user_turns - n + if surface == "cli": + from hermes_cli.cli_session_mixin import CLISessionMixin + cli = CLISessionMixin.__new__(CLISessionMixin) + cli._session_db, cli.session_id, cli.conversation_history, cli.agent = db, sid, warm, None + cli._prefill_input_buffer = MagicMock() + cli.undo_last(n) + return cli.conversation_history if len(cli.conversation_history) < len(warm) else None + from tui_gateway.server import _rewind_active_session_history + session = {"agent": SimpleNamespace(_session_messages=warm), "history": list(warm), + "history_lock": threading.Lock(), "history_version": 0, "session_key": sid} + return _rewind_active_session_history(session, ordinal) + + +@pytest.fixture() +def db(tmp_path, monkeypatch): + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + from tui_gateway import server + handle = SessionDB(db_path=tmp_path / "state.db") + monkeypatch.setattr(server, "_db", handle) + yield handle + handle.close() + + +@pytest.mark.parametrize("n", [1, 2]) +def test_every_surface_persists_the_same_active_set(db, n): + expected = None + for surface in SURFACES: + sid = f"rewind-{surface}-{n}" + _seed(db, sid) + assert _rewind_via(surface, db, sid, n) + rows = [(role, content, active) for _id, role, content, active in _active_rows(db, sid)] + assert rows == (expected := expected or rows), surface + assert [c for _r, c, a in expected if a] == [f"q{i}" if k == 0 else f"a{i}" + for i in range(1, 4 - n) for k in (0, 1)] + + +@pytest.mark.parametrize("surface", SURFACES) +def test_out_of_range_target_changes_nothing_on_every_surface(db, surface): + sid = f"oob-{surface}" + _seed(db, sid, turns=1) + before = _active_rows(db, sid) + if surface == "tui": + with pytest.raises(ValueError): # the TUI's own error class for a vanished target + _rewind_via(surface, db, sid, 2) + else: + # CLI /undo N and gateway /undo N clamp to the oldest turn, so their out-of-range case is a + # transcript with no user turn at all: both refuse without touching the store. + db.create_session(sid + "-empty", source="cli") + if surface == "gateway": + assert _rewind_via(surface, db, sid + "-empty", 1) is None + else: + with pytest.raises(RewindTargetUnavailableError): + db.rewind_user_turn(sid + "-empty", -1, warm_history=[]) + assert _active_rows(db, sid + "-empty") == [] + assert _active_rows(db, sid) == before diff --git a/tui_gateway/session_workdir.py b/tui_gateway/session_workdir.py index a10b8b54fd..2418788824 100644 --- a/tui_gateway/session_workdir.py +++ b/tui_gateway/session_workdir.py @@ -370,84 +370,28 @@ def _session_db(session: dict): def _rewind_active_session_history( session: dict, user_ordinal: int, *, require_retryable: bool = False) -> tuple[list[dict], dict, int]: """Rewind one canonical user turn while retaining carrier scaffolding. Caller holds ``history_lock``. Persistent - sessions archive the target and tail, inserting a composite carrier's hidden handoff in the same transaction; memory - is installed only after the durable commit, from the validated prefix + returned scaffold row id (no reload).""" - from agent.context_compressor import ( - history_before_user_originated_turn, retryable_user_text, split_user_originated_turn, - user_originated_turn_view) - from agent.memory_manager import sanitize_context - from agent.tool_dispatch_helpers import _is_multimodal_tool_result, _multimodal_text_summary - - def _comparison_content(message: dict) -> Any: - content = message.get("content") - if _is_multimodal_tool_result(content): - content = _multimodal_text_summary(content) - elif isinstance(content, list): - text_parts = [ - str(part.get("text", "")) if part.get("type") == "text" else "[screenshot]" - for part in content - if isinstance(part, dict) and part.get("type") in {"text", "image", "image_url", "input_image"}] - content = "\n".join(text_parts) if text_parts else None - if message.get("role") in {"user", "assistant"} and isinstance(content, str): - return sanitize_context(content).strip() - return content - - def _user_indices(messages: list[dict]) -> list[int]: - return [i for i, m in enumerate(messages) if user_originated_turn_view(m) is not None] + sessions go through ``SessionDB.rewind_user_turn`` (the durable transcript is the authority; memory is installed + only after the commit); a session without a key rewinds the warm history alone.""" + from agent.context_compressor import history_before_user_originated_turn, retryable_user_text, user_originated_turn_view history = _history_without_ephemeral_scaffolding(session.get("history", [])) - user_indices = _user_indices(history) + user_indices = [i for i, m in enumerate(history) if user_originated_turn_view(m) is not None] if user_ordinal < 0 or user_ordinal >= len(user_indices): raise ValueError("target user message is no longer in session history") - target_index = user_indices[user_ordinal] - installed, live_view = history_before_user_originated_turn(history, target_index) - rewound_count = len(history) - target_index - session_key = str(session.get("session_key") or "").strip() - persisted = False if session_key: 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) - durable_user_indices = _user_indices(durable) - if len(durable_user_indices) != len(user_indices): - raise RuntimeError("session history changed before the rewind could be persisted") - durable_target_index = durable_user_indices[user_ordinal] - durable_target = durable[durable_target_index] - durable_prefix, durable_live_view = history_before_user_originated_turn(durable, durable_target_index) - if _comparison_content(durable_live_view) != _comparison_content(live_view): - raise RuntimeError("session history changed before the rewind could be persisted") - if not isinstance(target_row_id := durable_target.get("_row_id"), int): - raise RuntimeError("rewind target has no durable row identity") - 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, target_row_id, preserve_compaction_handoff=scaffold is not None, - expected_active_ids=expected_active_ids, expected_target_content=durable_live_view.get("content")) - if scaffold is not None: - if not isinstance(replacement_id := result.get("replacement_message_id"), int): - raise RuntimeError("rewind commit did not return the replacement scaffold id") - durable_prefix[-1].update(_row_id=replacement_id, _db_persisted=True) - installed[-1] = durable_prefix[-1] - # Clients address follow-ups by durable row id: keep the richer warm content but copy row identities when - # the shapes align. - if len(installed) == len(durable_prefix) and all( - warm.get("role") == durable_message.get("role") - and bool(warm.get("display_kind")) == bool(durable_message.get("display_kind")) - and _comparison_content(warm) == _comparison_content(durable_message) - for warm, durable_message in zip(installed, durable_prefix) - ): - for warm, durable_message in zip(installed, durable_prefix): - if isinstance(row_id := durable_message.get("_row_id"), int): - warm["_row_id"] = row_id - 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")) + outcome = db.rewind_user_turn( + session_key, user_ordinal, warm_history=history, require_retryable=require_retryable) + installed, live_view, rewound_count = outcome.prefix, outcome.live_view, outcome.rewound_count + else: + target_index = user_indices[user_ordinal] + installed, live_view = history_before_user_originated_turn(history, target_index) + rewound_count = len(history) - target_index + if require_retryable: + retryable_user_text(live_view.get("content")) installed = [message.copy() for message in installed] session["history"] = installed @@ -456,9 +400,9 @@ def _rewind_active_session_history( if agent is not None: agent._session_messages = installed if hasattr(agent, "_last_flushed_db_idx"): - agent._last_flushed_db_idx = len(installed) if persisted else 0 + agent._last_flushed_db_idx = len(installed) if session_key else 0 if hasattr(agent, "_db_flush_scan_prefix"): - agent._db_flush_scan_prefix = installed[:] if persisted else None + agent._db_flush_scan_prefix = installed[:] if session_key else None return installed, live_view, rewound_count