diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index a2248e9b63..94694db158 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -739,6 +739,19 @@ def repair_message_sequence(agent, messages: List[Dict]) -> int: and merged[-1].get("role") == "user" ): prev = merged[-1] + # A summary carrier followed by a new user row is a deliberate + # durable shape after retry/rewind. Do not absorb the fresh ask + # into the already-persisted carrier: mutating that dict can make + # the only in-memory copy diverge from its durable row. Provider + # sanitizers merge copies later when strict alternation requires + # it, without rewriting either durable message. + from agent.context_compressor import split_user_originated_turn + + handoff, _ = split_user_originated_turn(prev) + if handoff is not None: + merged.append(msg) + continue + prev_content = prev.get("content", "") new_content = msg.get("content", "") # Only merge plain-text content; leave multimodal (list) @@ -1395,8 +1408,6 @@ def drop_thinking_only_and_merge_users( ) ] dropped = len(messages) - len(kept) - if dropped == 0: - return messages # Pass 2: merge any newly-adjacent user messages. merged: List[Dict[str, Any]] = [] @@ -1446,6 +1457,9 @@ def drop_thinking_only_and_merge_users( else: merged.append(m) + if dropped == 0 and merges == 0: + return messages + _ra().logger.debug( "Pre-call sanitizer: dropped %d thinking-only assistant turn(s), " "merged %d adjacent user message(s)", diff --git a/agent/context_compressor.py b/agent/context_compressor.py index 3eb05b6b97..290016d193 100644 --- a/agent/context_compressor.py +++ b/agent/context_compressor.py @@ -16,6 +16,7 @@ Improvements over v2: - Richer tool call/result detail in summarizer input """ +import copy import hashlib import json import logging @@ -7294,6 +7295,218 @@ def is_compaction_summary_message(message: Any) -> bool: return ContextCompressor._is_context_summary_content(content) +# Display metadata that describes the durable message independently of the +# compaction wrapper. Other metadata may describe a synthetic timeline event +# and must not make that event look human after projection. +SUMMARY_CARRIER_DURABLE_DISPLAY_METADATA_KEYS = ("reactions",) + + +def _handoff_only_content(content: Any) -> Any: + """Project summary-bearing content to the synthetic handoff alone. + + The compressor has two composite layouts. Ordinary merge-into-tail keeps + the live content before ``_MERGED_SUMMARY_DELIMITER``; the force-user- + leading layout keeps it after ``_SUMMARY_END_MARKER``. This is the inverse + of ``_strip_context_summary_handoff_message`` and deliberately never keeps + live media blocks. + """ + if isinstance(content, str): + if _MERGED_SUMMARY_DELIMITER in content: + suffix = content.split(_MERGED_SUMMARY_DELIMITER, 1)[1].lstrip() + marker_idx = suffix.find(_SUMMARY_END_MARKER) + if marker_idx >= 0: + return suffix[: marker_idx + len(_SUMMARY_END_MARKER)] + return suffix + marker_idx = content.find(_SUMMARY_END_MARKER) + if marker_idx >= 0: + return content[: marker_idx + len(_SUMMARY_END_MARKER)] + return content + + if not isinstance(content, list): + return content + + # Ordinary merge: the summary suffix begins in the delimiter-bearing text + # part. Do not retain later parts: malformed/legacy rows may carry live + # media there rather than synthetic scaffold content. + for item in content: + text = ( + item + if isinstance(item, str) + else item.get("text") + if isinstance(item, dict) + else None + ) + if not isinstance(text, str) or _MERGED_SUMMARY_DELIMITER not in text: + continue + suffix = text.split(_MERGED_SUMMARY_DELIMITER, 1)[1].lstrip() + marker_idx = suffix.find(_SUMMARY_END_MARKER) + if marker_idx >= 0: + suffix = suffix[: marker_idx + len(_SUMMARY_END_MARKER)] + if not suffix: + return [] + if isinstance(item, dict): + copied = item.copy() + copied["text"] = suffix + return [copied] + return [suffix] + + # Force-user-leading merge: keep textual parts through the end marker and + # truncate the marker-bearing part before the live ask. + projected: list[Any] = [] + for item in content: + text = ( + item + if isinstance(item, str) + else item.get("text") + if isinstance(item, dict) + else None + ) + if isinstance(text, str) and _SUMMARY_END_MARKER in text: + prefix = text.split(_SUMMARY_END_MARKER, 1)[0] + _SUMMARY_END_MARKER + if isinstance(item, dict): + copied = item.copy() + copied["text"] = prefix + projected.append(copied) + else: + projected.append(prefix) + return projected + if isinstance(text, str): + projected.append(item.copy() if isinstance(item, dict) else item) + return projected + + +def split_user_originated_turn( + message: Any, +) -> tuple[Optional[Dict[str, Any]], Optional[Dict[str, Any]]]: + """Split a user row into hidden handoff scaffold and canonical live view. + + Returns ``(handoff_only, live_view)``. A normal human row has no + handoff; a pure compaction handoff has no live view; a composite carrier + has both. Rewritten projections are fresh dictionaries and never retain + stale API-content or physical persistence identity. + """ + if not isinstance(message, dict) or message.get("role") != "user": + return None, None + + is_summary = is_compaction_summary_message(message) + handoff: Optional[Dict[str, Any]] = None + if is_summary: + handoff = { + "role": "user", + "content": _handoff_only_content(message.get("content")), + COMPRESSED_SUMMARY_METADATA_KEY: True, + "display_kind": "hidden", + } + if COMPRESSED_SUMMARY_HAS_USER_TURN_KEY in message: + handoff[COMPRESSED_SUMMARY_HAS_USER_TURN_KEY] = bool( + message.get(COMPRESSED_SUMMARY_HAS_USER_TURN_KEY) + ) + if message.get(MICRO_COMPACT_MARKER_KEY): + handoff[MICRO_COMPACT_MARKER_KEY] = True + if message.get("timestamp") is not None: + handoff["timestamp"] = message["timestamp"] + drop_stale_api_content(handoff) + + # Hidden is the legacy physical wrapper used for compaction rows and + # does not hide a successfully unwrapped human payload. Other typed + # display rows are synthetic timeline events, never user input. + display_kind = message.get("display_kind") + if display_kind and display_kind != "hidden": + return handoff, None + candidate = ContextCompressor._strip_context_summary_handoff_message(message) + if candidate is None: + return handoff, None + else: + if message.get("display_kind"): + return None, None + candidate = message.copy() + + candidate.pop(COMPRESSED_SUMMARY_METADATA_KEY, None) + candidate.pop(COMPRESSED_SUMMARY_HAS_USER_TURN_KEY, None) + candidate.pop(MICRO_COMPACT_MARKER_KEY, None) + candidate.pop(_DB_PERSISTED_MARKER, None) + if is_summary: + candidate.pop("_row_id", None) + candidate.pop("display_kind", None) + candidate.pop("display_metadata", None) + carrier_metadata = message.get("display_metadata") + if isinstance(carrier_metadata, dict): + durable_metadata = { + key: copy.deepcopy(carrier_metadata[key]) + for key in SUMMARY_CARRIER_DURABLE_DISPLAY_METADATA_KEYS + if key in carrier_metadata + } + if durable_metadata: + candidate["display_metadata"] = durable_metadata + drop_stale_api_content(candidate) + if ContextCompressor._is_synthetic_compression_user_turn(candidate): + return handoff, None + if not ContextCompressor._is_actionable_user_turn(candidate): + return handoff, None + return handoff, candidate + + +def user_originated_turn_view(message: Any) -> Optional[Dict[str, Any]]: + """Return the live human-authored projection of a user row, if any.""" + return split_user_originated_turn(message)[1] + + +def history_before_user_originated_turn( + messages: List[Dict[str, Any]], + index: int, +) -> tuple[List[Dict[str, Any]], Dict[str, Any]]: + """Return a rewind prefix and canonical live view for ``index``. + + When the selected row is a composite carrier, the hidden handoff scaffold + remains at the new history head while the live ask and later rows are + removed. This retains the only representation of already-compacted turns. + """ + if index < 0 or index >= len(messages): + raise IndexError("user turn index is outside the transcript") + handoff, live_view = split_user_originated_turn(messages[index]) + if live_view is None: + raise ValueError("selected row is not a user-originated turn") + prefix = [message.copy() for message in messages[:index]] + if handoff is not None: + prefix.append(handoff) + return prefix, live_view + + +def retryable_user_text(content: Any) -> str: + """Return lossless retry text or raise before destructive mutation. + + Retry has no attachment replay protocol. Media and unknown structured + parts therefore fail closed; already-persisted strings are replayed as + text, including any textual degradation labels. Structured content is + flattened only when every part is plain text. + """ + if isinstance(content, str): + text = content + elif isinstance(content, list): + chunks: list[str] = [] + for part in content: + if isinstance(part, str): + chunks.append(part) + continue + if not isinstance(part, dict): + raise ValueError("retry does not support non-text content") + if part.get("type") not in {"text", "input_text", "output_text"}: + raise ValueError("retry does not support media or unknown content parts") + if set(part) - {"type", "text"}: + raise ValueError("retry cannot losslessly flatten annotated text parts") + part_text = part.get("text") + if not isinstance(part_text, str): + raise ValueError("retry text parts must contain text") + chunks.append(part_text) + text = "".join(chunks) + else: + raise ValueError("retry does not support non-text content") + + if not text.strip(): + raise ValueError("retry found no text to send") + return text + + def _handoff_carries_live_user_content(message: Any) -> bool: """Return True when a summary-bearing row still carries a live user ask. @@ -7302,21 +7515,16 @@ def _handoff_carries_live_user_content(message: Any) -> bool: ask, leaving a non-empty remainder after ``_SUMMARY_END_MARKER``. Either shape must remain actionable (#80622 must not treat them as sole-handoff). - Delegates to ``_strip_context_summary_handoff_message`` — the canonical - "does anything survive once the handoff is removed" logic (it also - handles multimodal list content and returns ``None`` for a merged-shaped - row whose preserved prior tail is EMPTY, which a bare - ``classify_summary_content(...) == "merged"`` check would wrongly treat - as live). Callers must pre-filter with ``is_compaction_summary_message``: - for non-summary rows the strip helper returns the message unchanged, - which would read as "carries live content" here. + Delegates to the canonical live-user projection. That projection uses + ``_strip_context_summary_handoff_message`` for string and multimodal + carriers, then excludes typed synthetic rows and empty/non-actionable + remainders. Callers must still pre-filter with + ``is_compaction_summary_message`` because an ordinary human row naturally + has a live-user projection too. """ if not isinstance(message, dict): return False - return ( - ContextCompressor._strip_context_summary_handoff_message(message) - is not None - ) + return user_originated_turn_view(message) is not None def reference_handoff_would_drive_next_model_call( @@ -7372,15 +7580,8 @@ def is_user_originated_turn(message: Any) -> bool: this instead of ``role == "user" and not display_kind`` — standalone handoffs with ``_compressed_summary_has_user_turn`` were previously left without ``display_kind=hidden`` and could be mistaken for real asks (#80622). - Summary-bearing rows are never user-originated, even when they embed a - live ask after the end marker (callers that need that text should unwrap). + Summary-bearing rows count only when their canonical live-user projection + recovers an actionable ask. Pure handoffs and typed synthetic rows never + count. """ - if not isinstance(message, dict) or message.get("role") != "user": - return False - if message.get("display_kind"): - return False - if is_compaction_summary_message(message): - return False - if ContextCompressor._is_synthetic_compression_user_turn(message): - return False - return ContextCompressor._is_actionable_user_turn(message) + return user_originated_turn_view(message) is not None diff --git a/agent/turn_context.py b/agent/turn_context.py index dfa5fbbd8e..eef1ecc8d8 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -273,18 +273,26 @@ def reanchor_current_turn_user_idx(messages: List[Any], user_message: Any) -> in scaffolding, not the active ask. Returns -1 when the list has no user-originated message at all. """ - from agent.context_compressor import is_user_originated_turn + from agent.context_compressor import user_originated_turn_view fallback = -1 for i in range(len(messages) - 1, -1, -1): msg = messages[i] if not (isinstance(msg, dict) and msg.get("role") == "user"): continue + # Typed synthetic current events still need their physical persistence + # anchor when their raw content is unchanged. They are not eligible + # for the human-only fallback below. if msg.get("content") == user_message: return i + live_view = user_originated_turn_view(msg) + if live_view is None: + continue + if live_view.get("content") == user_message: + return i # Prefer a real human turn over a synthetic handoff / continuation # marker when the exact content was rewritten by merge-into-tail. - if fallback < 0 and is_user_originated_turn(msg): + if fallback < 0: fallback = i return fallback diff --git a/cli.py b/cli.py index 2e822be374..cd218ab1c8 100644 --- a/cli.py +++ b/cli.py @@ -8710,6 +8710,121 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): print(f" Resume the live session with: hermes --resume {self.session_id}") 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 run_agent 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 + + durable = self._session_db.get_messages_as_conversation( + self.session_id, + include_row_ids=True, + ) + warm_persistence_history = [ + message + for message in warm_history + if not _is_ephemeral_scaffolding(message) + ] + warm_user_indices = [ + index + for index, message in enumerate(warm_persistence_history) + if user_originated_turn_view(message) is not None + ] + durable_user_indices = [ + index + for index, message in enumerate(durable) + if user_originated_turn_view(message) is not None + ] + if len(durable_user_indices) != len(warm_user_indices): + raise RuntimeError( + "session history changed before the rewind could be persisted" + ) + 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 RuntimeError( + "session history changed before the rewind could be persisted" + ) + 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, + 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 retry_last(self): """Retry the last user message by removing the last exchange and re-sending. @@ -8727,22 +8842,68 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): # CLI resume counting and list_recent_user_messages. Compaction # handoffs are excluded too (durable role=user, sometimes without # display_kind on legacy sessions; #80622). - from agent.context_compressor import is_user_originated_turn + from agent.context_compressor import ( + history_before_user_originated_turn, + retryable_user_text, + user_originated_turn_view, + ) + from agent.memory_manager import sanitize_context + from run_agent import _is_ephemeral_scaffolding - last_user_idx = None - for i in range(len(self.conversation_history) - 1, -1, -1): - msg = self.conversation_history[i] - if is_user_originated_turn(msg): - last_user_idx = i - break + warm_history = list(self.conversation_history) + + user_indices = [ + index + for index, message in enumerate(warm_history) + if not _is_ephemeral_scaffolding(message) + and user_originated_turn_view(message) is not None + ] - if last_user_idx is None: + if not user_indices: print("(._.) No user message found to retry.") return None + last_user_idx = user_indices[-1] - # Extract the message text and remove everything from that point forward - last_message = self.conversation_history[last_user_idx].get("content", "") - self.conversation_history = self.conversation_history[:last_user_idx] + # Resolve a lossless live payload before touching either persistence or + # memory. A force-user-leading compaction row is one physical carrier: + # its historical handoff remains in the prefix while only the embedded + # human ask is retried. Media cannot be replayed by /retry, so fail + # closed before archiving anything. + try: + truncated, live_view = history_before_user_originated_turn( + warm_history, last_user_idx + ) + live_content = live_view.get("content") + if isinstance(live_content, str): + live_content = sanitize_context(live_content).strip() + last_message = retryable_user_text(live_content) + except ValueError as exc: + print(f"(._.) Cannot retry that message safely: {exc}") + return None + + # Persist the rewind before publishing the shorter in-memory view. + # The DB owns the physical carrier split so the archived original and + # retained scaffold are committed atomically. A plain user row keeps + # the legacy rewind shape (no replacement scaffold). + 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, + ) + except Exception as exc: + print(f"(x_x) Retry rewind failed; history was not changed: {exc}") + return None + + self.conversation_history = truncated + if self.agent is not None: + if hasattr(self.agent, "_session_messages"): + self.agent._session_messages = self.conversation_history + if hasattr(self.agent, "_last_flushed_db_idx"): + self.agent._last_flushed_db_idx = len(self.conversation_history) + if hasattr(self.agent, "_db_flush_scan_prefix"): + self.agent._db_flush_scan_prefix = self.conversation_history[:] print(f"(^_^)b Retrying: \"{last_message[:60]}{'...' if len(last_message) > 60 else ''}\"") return last_message @@ -8781,59 +8942,62 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): # messages (exclude display_kind timeline rows and compaction # handoffs — same predicate as list_recent_user_messages, resume # turn counting, and /retry; #80622). - from agent.context_compressor import is_user_originated_turn + from agent.context_compressor import ( + history_before_user_originated_turn, + user_originated_turn_view, + ) + from run_agent import _is_ephemeral_scaffolding - user_indices = [] - for i in range(len(self.conversation_history) - 1, -1, -1): - msg = self.conversation_history[i] - if is_user_originated_turn(msg): - user_indices.append(i) - if len(user_indices) >= n: - break + warm_history = list(self.conversation_history) + + user_indices = [ + index + for index, message in enumerate(warm_history) + if not _is_ephemeral_scaffolding(message) + and user_originated_turn_view(message) is not None + ] if not user_indices: print("(._.) No user message found to undo.") return - # The oldest of the collected user messages is our truncation point. - cut_idx = user_indices[-1] - turns_undone = len(user_indices) + turns_undone = min(n, len(user_indices)) + target_ordinal = len(user_indices) - turns_undone + cut_idx = user_indices[target_ordinal] - removed_count = len(self.conversation_history) - cut_idx - removed_msg = self.conversation_history[cut_idx].get("content", "") - removed_text = self._undo_content_to_text(removed_msg) - - # Truncate the in-memory history to before that user message. - self.conversation_history = self.conversation_history[:cut_idx] + removed_count = len(warm_history) - cut_idx + truncated, live_view = history_before_user_originated_turn( + warm_history, cut_idx + ) + removed_text = self._undo_content_to_text(live_view.get("content")) # Soft-delete the truncated rows on disk so re-prompts and search # see the clean transcript while the rows survive for audit. rewound_rows = 0 if self._session_db is not None and self.session_id: try: - recents = self._session_db.list_recent_user_messages( - self.session_id, limit=max(turns_undone, 10) + truncated, durable_live_view, result = ( + self._rewind_persisted_user_turn( + warm_history=warm_history, + user_ordinal=target_ordinal, + warm_live_view=live_view, + ) ) - if recents: - target_idx = min(turns_undone - 1, len(recents) - 1) - target_id = recents[target_idx]["id"] - result = self._session_db.rewind_to_message( - self.session_id, target_id - ) - rewound_rows = result.get("rewound_count", 0) - # Prefer the DB's decoded target text for the prefill — - # it's the canonical persisted copy. - db_text = self._undo_content_to_text( - (result.get("target_message") or {}).get("content") - ) - if db_text: - removed_text = db_text - except ValueError as e: - # Non-user target / cross-session — keep the in-memory undo - # but skip the soft-delete; surface a debug-level note. - logger.debug("undo: soft-delete skipped: %s", e) + # Canonicalize the editable prefill before mutation. The raw + # physical carrier contains 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) except Exception as e: - logger.debug("undo: soft-delete failed: %s", e) + logger.debug("undo: durable rewind failed: %s", e) + print(f"(x_x) Undo failed; history was not changed: {e}") + return + + # Publish only after the durable rewind succeeds (or no store exists). + self.conversation_history = truncated # Agent surgery: invalidate the system-prompt cache and reset the # flush index so the next turn re-flushes from the truncated head. @@ -8848,6 +9012,10 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): self.agent._last_flushed_db_idx = len(self.conversation_history) except Exception: pass + if hasattr(self.agent, "_session_messages"): + self.agent._session_messages = self.conversation_history + if hasattr(self.agent, "_db_flush_scan_prefix"): + self.agent._db_flush_scan_prefix = self.conversation_history[:] # Notify memory providers — same hook /branch fires, with the # rewound flag so per-turn document caches invalidate (#6672, #21910). try: diff --git a/gateway/session.py b/gateway/session.py index 1d4db2f5ac..21228473e6 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -3874,7 +3874,13 @@ class SessionStore: ) return [] - def rewind_session(self, session_id: str, n: int = 1) -> Optional[Dict[str, Any]]: + def rewind_session( + self, + session_id: str, + n: int = 1, + *, + require_retryable_composite: bool = False, + ) -> Optional[Dict[str, Any]]: """Back up ``n`` user turns via soft-delete, keeping rows for audit. Unlike :meth:`rewrite_transcript` (a hard replace used by /retry), @@ -3885,45 +3891,86 @@ class SessionStore: Returns a dict ``{"rewound_count", "turns_undone", "target_text"}`` on success, or ``None`` if there's no DB or no user message to back up to. ``n`` clamps to the oldest user turn when it exceeds the turn count. + ``require_retryable_composite`` is the gateway ``/retry`` guard: the + selected current turn must still be a composite carrier, and its live + payload must be losslessly replayable as text before anything changes. """ if not self._db: return None self._clear_dirty_transcript(session_id) if n < 1: n = 1 + from agent.context_compressor import ( + retryable_user_text, + split_user_originated_turn, + user_originated_turn_view, + ) + try: - recents = self._db.list_recent_user_messages(session_id, limit=max(n, 10)) + 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) + 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: + if handoff is None: + return None + target_text = retryable_user_text(target_view.get("content")) except Exception as e: - logger.debug("rewind_session: failed to list user messages: %s", e) + logger.debug("rewind_session: failed to resolve canonical target: %s", e) return None - if not recents: - return None - target_idx = min(n - 1, len(recents) - 1) - target_id = recents[target_idx]["id"] try: - result = self._db.rewind_to_message(session_id, target_id) + result = self._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 as e: logger.debug("rewind_session: %s", e) return None except Exception as e: logger.debug("rewind_session: rewind_to_message failed: %s", e) return None - target_msg = result.get("target_message") or {} - content = target_msg.get("content") or "" - if isinstance(content, list): - parts = [ - p.get("text", "") - for p in content - if isinstance(p, dict) and p.get("type") == "text" - ] - target_text = "\n".join(t for t in parts if t) - elif isinstance(content, str): - target_text = content - else: - target_text = "" + # ``target_view`` is the canonical live projection of the physical DB + # row. For a composite carrier, the raw target contains the historical + # summary wrapper and must never be echoed back as the editable prompt. + if not require_retryable_composite: + content = target_view.get("content") or "" + if isinstance(content, list): + parts = [ + p.get("text", "") + for p in content + if isinstance(p, dict) and p.get("type") == "text" + ] + target_text = "\n".join(t for t in parts if t) + elif isinstance(content, str): + target_text = content + else: + target_text = "" return { "rewound_count": result.get("rewound_count", 0), - "turns_undone": target_idx + 1, + "turns_undone": turns_undone, "target_text": target_text, } diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index e49891749b..7f72f0df1f 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -2597,34 +2597,63 @@ class GatewaySlashCommandsMixin: # auto_continue / hidden); clients never count them as user turns. # Without this filter /retry rewrote the transcript around a marker # and re-sent opaque bookkeeping text (same class as the TUI ordinal). - last_user_msg = None last_user_idx = None - # is_user_originated_turn: excludes display_kind bookkeeping AND - # compaction handoffs (durable role=user, sometimes without - # display_kind on legacy sessions; #80622) — /retry must never - # re-send a reference-only summary as if the user asked it. - from agent.context_compressor import is_user_originated_turn + # The canonical projection excludes bookkeeping and pure handoffs while + # still recognizing a real ask embedded in a compaction carrier. + from agent.context_compressor import ( + history_before_user_originated_turn, + retryable_user_text, + split_user_originated_turn, + user_originated_turn_view, + ) for i in range(len(history) - 1, -1, -1): msg = history[i] - if is_user_originated_turn(msg): - last_user_msg = msg.get("content", "") + if user_originated_turn_view(msg) is not None: last_user_idx = i break - - if not last_user_msg: + + if last_user_idx is None: return t("gateway.retry.no_previous") - - # Truncate history to before the last user message and persist only the - # live view. After in-place compaction the pre-compaction transcript - # lives on as active=0/compacted=1 rows under this same session id, and - # a bare rewrite (active_only=False) would DELETE them (same class as - # #61145). /retry never intends to purge archived history, so avoid a - # separate existence probe: it could fail open or race with the write. - truncated = history[:last_user_idx] - await self.async_session_store.rewrite_transcript( - session_entry.session_id, truncated, active_only=True - ) + + # Resolve the live text and the scaffold-preserving prefix before any + # transcript write. Messaging retries cannot reconstruct attachments; + # reject media/unknown content without truncating the session. + try: + truncated, live_view = history_before_user_originated_turn( + history, last_user_idx + ) + last_user_msg = retryable_user_text(live_view.get("content")) + handoff, _ = split_user_originated_turn(history[last_user_idx]) + except ValueError as exc: + return f"Cannot retry that message safely: {exc}" + + if handoff is not None: + # A composite carrier is one physical row containing both the + # retained summary and the live ask. Let the carrier-aware rewind + # 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, + ) + if rewind_result is None: + return "Retry failed; transcript was not changed." + # The store reselects and validates the latest carrier on the same + # snapshot used by the atomic rewind. A concurrent newer turn can + # therefore never be removed while this handler resends stale text. + last_user_msg = rewind_result["target_text"] + else: + # After in-place compaction the pre-compaction transcript lives on + # as active=0/compacted=1 rows under this session id. active_only + # preserves that archive; a separate existence probe could fail + # open or race with the write. + if not await self.async_session_store.rewrite_transcript( + session_entry.session_id, truncated, active_only=True + ): + return "Retry failed; transcript was not changed." # Reset stored token count — transcript was truncated session_entry.last_prompt_tokens = 0 diff --git a/hermes_state.py b/hermes_state.py index 800a7828e1..9891014d6f 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -9473,8 +9473,37 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # Rewind (soft-delete) — see /rewind slash command + issue #21910 # ========================================================================= + @staticmethod + def _active_transcript_counts(conn, session_id: str) -> tuple[int, int]: + """Return active message/tool-call counts inside the caller's txn.""" + 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 + def rewind_to_message( - self, session_id: str, target_message_id: int + self, + session_id: str, + target_message_id: int, + *, + preserve_compaction_handoff: bool = False, + expected_active_ids: Optional[List[int]] = None, + expected_target_content: Any = None, ) -> Dict[str, Any]: """Soft-delete all messages with id >= ``target_message_id`` in *session_id*. @@ -9493,7 +9522,17 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) } Raises ``ValueError`` if the target message does not exist in - *session_id* or if its role is not ``"user"``. + *session_id* or if its role is not ``"user"``. With + ``preserve_compaction_handoff=True``, a composite summary carrier is + split inside the same write transaction: its original row is archived + and its canonical hidden handoff scaffold is inserted as the new head. + That opt-in result also contains ``replacement_message_id``. + + ``expected_active_ids`` optionally pins the ordered active row set. + ``expected_target_content`` additionally pins the selected canonical + live-user payload. Both checks run inside the write transaction before + any row or counter mutation. Presentation-only metadata changes (for + example Desktop reactions) deliberately do not invalidate a rewind. Always increments ``sessions.rewind_count`` — even when the target is already inactive — so the counter accurately reflects @@ -9502,29 +9541,71 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) target is a no-op on row state but still bumps the counter. """ - # 1) Validate target up-front (read-only, outside the write txn). - with self._lock: - row = self._conn.execute( + def _do(conn): + # Rewind changes the active transcript and must honor the same + # compression ownership/closed-parent guards as append writers. + self._check_transcript_write_guards(conn, session_id, None) + + if expected_active_ids is not None: + active_rows = conn.execute( + "SELECT id FROM messages " + "WHERE session_id = ? AND active = 1 ORDER BY id", + (session_id,), + ).fetchall() + active_ids = [int(active_row[0]) for active_row in active_rows] + if active_ids != 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), ).fetchone() - if row is None: - raise ValueError( - f"message {target_message_id} not found in session {session_id}" - ) - target_row = dict(row) - if target_row.get("role") != "user": - raise ValueError( - f"rewind target must be a 'user' message (got role=" - f"{target_row.get('role')!r}, id={target_message_id})" - ) + if row is None: + raise ValueError( + f"message {target_message_id} not found in session {session_id}" + ) + target_row = dict(row) + if target_row.get("role") != "user": + raise ValueError( + f"rewind target must be a 'user' message (got role=" + f"{target_row.get('role')!r}, id={target_message_id})" + ) - # Decode content for callers (prefill the prompt buffer). - target_row["content"] = self._decode_content(target_row.get("content")) + replacement_message_id: Optional[int] = None + replacement: Optional[Dict[str, Any]] = None + if preserve_compaction_handoff or expected_target_content is not None: + if not target_row.get("active"): + raise ValueError("rewind target is not active") + from agent.context_compressor import split_user_originated_turn - rewound: List[int] = [] + split_target = target_row.copy() + split_target["content"] = self._decode_content( + split_target.get("content") + ) + split_target["display_metadata"] = self._decode_display_metadata( + split_target.get("display_metadata") + ) + handoff, live_view = split_user_originated_turn(split_target) + if live_view is None: + raise ValueError("rewind target is not a user-originated turn") + live_content = live_view.get("content") + if isinstance(live_content, str): + live_content = sanitize_context(live_content).strip() + if ( + expected_target_content is not None + and live_content != expected_target_content + ): + raise RuntimeError( + "rewind target changed before it could be persisted" + ) + if preserve_compaction_handoff and handoff is None: + raise ValueError( + "preserve_compaction_handoff requires an active composite carrier" + ) + replacement = handoff if preserve_compaction_handoff else None - def _do(conn): cursor = conn.execute( "SELECT id FROM messages " "WHERE session_id = ? AND id >= ? AND active = 1", @@ -9537,28 +9618,48 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) f"UPDATE messages SET active = 0 WHERE id IN ({placeholders})", ids, ) + if replacement is not None: + self._insert_message_rows(conn, session_id, [replacement]) + inserted = conn.execute("SELECT last_insert_rowid()").fetchone() + replacement_message_id = int(inserted[0]) conn.execute( "UPDATE sessions SET rewind_count = COALESCE(rewind_count, 0) + 1 " "WHERE id = ?", (session_id,), ) - return ids - - rewound = self._execute_write(_do) - - # 2) Compute new head id (largest still-active row id in session). - with self._lock: - head_row = self._conn.execute( + message_count, tool_call_count = self._active_transcript_counts( + conn, session_id + ) + conn.execute( + "UPDATE sessions SET message_count = ?, tool_call_count = ? " + "WHERE id = ?", + (message_count, tool_call_count, session_id), + ) + 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 + 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, rewound, new_head_id, replacement_message_id = ( + self._execute_write(_do) + ) + + # Decode content for callers (prefill the prompt buffer) without a + # second fallible database operation after the transaction commits. + target_row["content"] = self._decode_content(target_row.get("content")) + + result = { "rewound_count": len(rewound), "target_message": target_row, "new_head_id": new_head_id, } + if preserve_compaction_handoff: + result["replacement_message_id"] = replacement_message_id + return result def restore_rewound(self, session_id: str, since_message_id: int) -> int: """Mark inactive messages with id >= *since_message_id* active again. diff --git a/run_agent.py b/run_agent.py index 12e7da647e..0f3bd873ab 100644 --- a/run_agent.py +++ b/run_agent.py @@ -162,6 +162,7 @@ from agent.usage_pricing import normalize_usage from agent.context_compressor import ( # noqa: F401 COMPRESSED_SUMMARY_METADATA_KEY, ContextCompressor, + user_originated_turn_view, ) from agent.retry_utils import jittered_backoff # noqa: F401 from agent.prompt_builder import ( # noqa: F401 # re-exported via _ra() / mock.patch("run_agent.") / from run_agent import @@ -2261,6 +2262,7 @@ class AIAgent: "hidden" if ( msg.get(COMPRESSED_SUMMARY_METADATA_KEY) + and user_originated_turn_view(msg) is None and ( ContextCompressor.classify_summary_content( msg.get("content") diff --git a/tests/agent/test_reference_handoff_active_turn.py b/tests/agent/test_reference_handoff_active_turn.py index 4e62b0f9b8..45df7c76b6 100644 --- a/tests/agent/test_reference_handoff_active_turn.py +++ b/tests/agent/test_reference_handoff_active_turn.py @@ -10,6 +10,8 @@ the already-completed work. from __future__ import annotations +import pytest + from agent.context_compressor import ( COMPRESSED_SUMMARY_HAS_USER_TURN_KEY, COMPRESSED_SUMMARY_METADATA_KEY, @@ -17,13 +19,18 @@ from agent.context_compressor import ( HISTORICAL_TASK_HEADING, SUMMARY_PREFIX, _SUMMARY_END_MARKER, + history_before_user_originated_turn, is_compaction_summary_message, is_user_originated_turn, reference_handoff_would_drive_next_model_call, + retryable_user_text, + split_user_originated_turn, + user_originated_turn_view, ) from agent.conversation_loop import ( _should_skip_model_call_for_reference_handoff, ) +from agent.agent_runtime_helpers import repair_message_sequence from agent.turn_context import reanchor_current_turn_user_idx @@ -39,6 +46,12 @@ def _standalone_handoff(task: str = "finish the already-done refactor") -> dict: } +def _composite_handoff(ask: str = "REAL ASK") -> dict: + handoff = _standalone_handoff() + handoff["content"] += f"\n\n{ask}" + return handoff + + class TestReferenceHandoffWouldDriveNextModelCall: def test_standalone_handoff_alone_drives(self): messages = [_standalone_handoff()] @@ -136,6 +149,52 @@ class TestUserOriginatedTurnPredicate: def test_plain_user_is_originated(self): assert is_user_originated_turn({"role": "user", "content": "hello"}) is True + def test_force_user_leading_carrier_is_originated_via_live_view(self): + carrier = _composite_handoff() + + handoff, live = split_user_originated_turn(carrier) + + assert is_user_originated_turn(carrier) is True + assert user_originated_turn_view(carrier)["content"] == "REAL ASK" + assert live["content"] == "REAL ASK" + assert handoff["display_kind"] == "hidden" + assert "REAL ASK" not in handoff["content"] + + def test_hidden_legacy_carrier_still_projects_but_typed_carrier_does_not(self): + hidden = {**_composite_handoff(), "display_kind": "hidden"} + typed = {**_composite_handoff(), "display_kind": "auto_continue"} + + assert user_originated_turn_view(hidden)["content"] == "REAL ASK" + assert user_originated_turn_view(typed) is None + + def test_rewind_prefix_preserves_only_the_carrier_scaffold(self): + messages = [ + {"role": "assistant", "content": "prefix"}, + _composite_handoff(), + {"role": "assistant", "content": "failed"}, + ] + + prefix, live = history_before_user_originated_turn(messages, 1) + + assert live["content"] == "REAL ASK" + assert prefix[0] == messages[0] + assert prefix[1]["display_kind"] == "hidden" + assert "REAL ASK" not in prefix[1]["content"] + assert messages[1]["content"].endswith("REAL ASK") + + def test_live_projection_keeps_reactions_but_drops_wrapper_metadata(self): + carrier = _composite_handoff() + carrier["display_metadata"] = { + "reactions": [{"emoji": "👍", "author": "user"}], + "synthetic_only": "do not project", + } + + live = user_originated_turn_view(carrier) + + assert live["display_metadata"] == { + "reactions": [{"emoji": "👍", "author": "user"}] + } + def test_display_kind_hidden_not_originated(self): assert ( is_user_originated_turn( @@ -172,6 +231,56 @@ class TestReanchorSkipsHandoffFallback: # Exact content rewritten by merge — fall back must not land on handoff. assert reanchor_current_turn_user_idx(messages, "rewritten ask") == 0 + def test_exact_live_projection_reanchors_to_composite_carrier(self): + messages = [ + {"role": "user", "content": "older ask"}, + _composite_handoff("rewritten ask"), + ] + + assert reanchor_current_turn_user_idx(messages, "rewritten ask") == 1 + + +class TestRetryableUserText: + def test_losslessly_flattens_text_only_parts(self): + assert retryable_user_text( + [ + {"type": "text", "text": "hello "}, + {"type": "input_text", "text": "world"}, + ] + ) == "hello world" + + def test_rejects_structured_media_parts(self): + with pytest.raises(ValueError, match="media or unknown"): + retryable_user_text( + [{"type": "image_url", "image_url": {"url": "data:image/png;base64,x"}}] + ) + + @pytest.mark.parametrize( + "text", + [ + "Explain why ![alt](https://example.test/a.png) renders", + "Review @file:README.md and this data:image/png example", + "Explain the literal [screenshot] marker", + "Compare [file: README.md] with [image|ybres:RID]", + ], + ) + def test_keeps_ordinary_text_that_uses_media_like_syntax(self, text): + assert retryable_user_text(text) == text + + +class TestCarrierAlternationRepair: + def test_fresh_user_does_not_mutate_persisted_carrier_dict(self): + carrier = _composite_handoff() + original = carrier.copy() + fresh = {"role": "user", "content": "NEXT ASK"} + messages = [carrier, fresh] + + assert repair_message_sequence(None, messages) == 0 + + assert messages == [carrier, fresh] + assert messages[0] is carrier + assert carrier == original + class TestNoToolCallsWithoutLaterRealUser: def test_historical_snapshot_alone_is_not_actionable(self): diff --git a/tests/cli/test_cli_retry.py b/tests/cli/test_cli_retry.py index b287b45754..8386184345 100644 --- a/tests/cli/test_cli_retry.py +++ b/tests/cli/test_cli_retry.py @@ -1,10 +1,42 @@ -"""Regression tests for CLI /retry history replacement semantics.""" +"""Regression tests for CLI /retry and carrier-aware rewind semantics.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from agent.context_compressor import ( + HISTORICAL_TASK_HEADING, + SUMMARY_PREFIX, + _SUMMARY_END_MARKER, +) +from hermes_state import SessionDB from tests.cli.test_cli_init import _make_cli +def _composite_carrier(ask="REAL ASK"): + return { + "role": "user", + "content": ( + f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n" + f"{_SUMMARY_END_MARKER}\n\n{ask}" + ), + } + + +def _message_rows(db, session_id): + rows = db._conn.execute( + "SELECT id, content, active FROM messages " + "WHERE session_id = ? ORDER BY id", + (session_id,), + ).fetchall() + return [tuple(row) for row in rows] + + def test_retry_last_truncates_history_before_requeueing_message(): cli = _make_cli() + cli._session_db = None cli.conversation_history = [ {"role": "user", "content": "first"}, {"role": "assistant", "content": "one"}, @@ -31,6 +63,7 @@ def test_retry_last_truncates_history_before_requeueing_message(): def test_process_command_retry_requeues_original_message_not_retry_command(): cli = _make_cli() + cli._session_db = None queued = [] class _Queue: @@ -47,3 +80,299 @@ def test_process_command_retry_requeues_original_message_not_retry_command(): assert queued == ["retry me"] assert cli.conversation_history == [] + + +def test_retry_fails_closed_when_warm_and_durable_targets_differ(tmp_path): + cli = _make_cli() + cli._session_db.close() + db = SessionDB(db_path=tmp_path / "state.db") + cli._session_db = db + cli.session_id = "cli-target-mismatch" + db.create_session(cli.session_id, source="cli") + db.append_message(cli.session_id, "user", "DURABLE ASK") + db.append_message(cli.session_id, "assistant", "old answer") + + history = [ + {"role": "user", "content": "WARM ASK"}, + {"role": "assistant", "content": "old answer"}, + ] + cli.conversation_history = history + cli._pending_input = MagicMock() + before_rows = _message_rows(db, cli.session_id) + + cli.process_command("/retry") + + cli._pending_input.put.assert_not_called() + assert cli.conversation_history is history + assert _message_rows(db, cli.session_id) == before_rows + db.close() + + +def test_retry_fails_closed_when_transcript_changes_after_snapshot( + tmp_path, monkeypatch +): + cli = _make_cli() + cli._session_db.close() + db = SessionDB(db_path=tmp_path / "state.db") + sibling = SessionDB(db_path=db.db_path) + cli._session_db = db + cli.session_id = "cli-cas-race" + db.create_session(cli.session_id, source="cli") + db.append_message(cli.session_id, "user", "RETRY ME") + db.append_message(cli.session_id, "assistant", "failed answer") + history = db.get_messages_as_conversation(cli.session_id) + cli.conversation_history = history + original_rewind = db.rewind_to_message + + def _append_then_rewind(*args, **kwargs): + sibling.append_message(cli.session_id, "assistant", "concurrent tail") + return original_rewind(*args, **kwargs) + + monkeypatch.setattr(db, "rewind_to_message", _append_then_rewind) + + assert cli.retry_last() is None + + assert cli.conversation_history is history + rows = db._conn.execute( + "SELECT content, active FROM messages " + "WHERE session_id = ? ORDER BY id", + (cli.session_id,), + ).fetchall() + assert [tuple(row) for row in rows] == [ + ("RETRY ME", 1), + ("failed answer", 1), + ("concurrent tail", 1), + ] + sibling.close() + db.close() + + +@pytest.mark.parametrize("command", ["retry", "undo"]) +def test_rewind_matches_warm_raw_carrier_to_durable_sanitized_sidecar( + tmp_path, command +): + from agent.memory_manager import sanitize_context + + cli = _make_cli() + cli._session_db.close() + db = SessionDB(db_path=tmp_path / "state.db") + cli._session_db = db + cli.session_id = f"cli-sanitized-{command}" + db.create_session(cli.session_id, source="cli") + raw_carrier = _composite_carrier( + " REAL ASK\n\n\nprivate\n " + )["content"] + db.append_message( + cli.session_id, + "user", + sanitize_context(raw_carrier).strip(), + api_content=raw_carrier, + ) + db.append_message(cli.session_id, "assistant", "failed answer") + durable = db.get_messages_as_conversation(cli.session_id) + cli.conversation_history = [ + {"role": "user", "content": raw_carrier}, + durable[1], + ] + cli._pending_input = MagicMock() + cli._prefill_input_buffer = MagicMock() + + if command == "retry": + cli.process_command("/retry") + cli._pending_input.put.assert_called_once_with("REAL ASK") + else: + cli.undo_last() + cli._prefill_input_buffer.assert_called_once_with("REAL ASK") + + assert len(cli.conversation_history) == 1 + scaffold = cli.conversation_history[0] + assert scaffold["display_kind"] == "hidden" + assert "REAL ASK" not in scaffold["content"] + active = db.get_messages_as_conversation(cli.session_id, include_row_ids=True) + assert active[0]["_row_id"] == scaffold["_row_id"] + assert active[0]["content"] == scaffold["content"] + db.close() + + +@pytest.mark.parametrize("command", ["retry", "undo"]) +@pytest.mark.parametrize("prefix_kind", ["buried_ephemeral", "old_media"]) +def test_rewind_keeps_the_richer_warm_prefix_after_validating_the_target( + tmp_path, command, prefix_kind +): + cli = _make_cli() + cli._session_db.close() + db = SessionDB(db_path=tmp_path / "state.db") + cli._session_db = db + cli.session_id = f"cli-projection-{prefix_kind}-{command}" + db.create_session(cli.session_id, source="cli") + + if prefix_kind == "buried_ephemeral": + history = [ + {"role": "user", "content": "OLDER ASK"}, + {"role": "assistant", "content": "candidate answer"}, + { + "role": "user", + "content": "[System: verify before stopping]", + "_verification_stop_synthetic": True, + }, + {"role": "assistant", "content": "verified answer"}, + {"role": "user", "content": "PLAIN TARGET"}, + {"role": "assistant", "content": "failed answer"}, + ] + durable_prefix = [ + ("user", "OLDER ASK"), + ("assistant", "candidate answer"), + ("assistant", "verified answer"), + ] + expected_prefix = ["OLDER ASK", "candidate answer", "verified answer"] + expected_active = [1, 1, 1, 0, 0] + else: + media_content = [ + {"type": "text", "text": "OLDER ASK"}, + {"type": "image_url", "image_url": {"url": "data:image/png,AA"}}, + ] + history = [ + {"role": "user", "content": media_content}, + {"role": "assistant", "content": "older answer"}, + {"role": "user", "content": "PLAIN TARGET"}, + {"role": "assistant", "content": "failed answer"}, + ] + durable_prefix = [ + ("user", "OLDER ASK\n[screenshot]"), + ("assistant", "older answer"), + ] + expected_prefix = [media_content, "older answer"] + expected_active = [1, 1, 0, 0] + + for role, content in durable_prefix: + db.append_message(cli.session_id, role, content) + db.append_message(cli.session_id, "user", "PLAIN TARGET") + db.append_message(cli.session_id, "assistant", "failed answer") + cli.conversation_history = history + cli._pending_input = MagicMock() + cli._prefill_input_buffer = MagicMock() + + if command == "retry": + cli.process_command("/retry") + cli._pending_input.put.assert_called_once_with("PLAIN TARGET") + else: + cli.undo_last() + cli._prefill_input_buffer.assert_called_once_with("PLAIN TARGET") + + assert [message.get("content") for message in cli.conversation_history] == ( + expected_prefix + ) + assert [row[2] for row in _message_rows(db, cli.session_id)] == expected_active + db.close() + + +def test_retry_last_durably_preserves_composite_carrier_scaffold(tmp_path): + cli = _make_cli() + cli._session_db.close() + db = SessionDB(db_path=tmp_path / "state.db") + cli._session_db = db + cli.session_id = "cli-carrier-retry" + db.create_session(cli.session_id, source="cli") + db.append_message(cli.session_id, "user", _composite_carrier()["content"]) + db.append_message(cli.session_id, "assistant", "failed answer") + cli.conversation_history = db.get_messages_as_conversation(cli.session_id) + old_history = cli.conversation_history + cli.agent = SimpleNamespace( + _session_messages=old_history, + _last_flushed_db_idx=len(old_history), + _db_flush_scan_prefix=list(old_history), + ) + + retry_msg = cli.retry_last() + + assert retry_msg == "REAL ASK" + assert len(cli.conversation_history) == 1 + scaffold = cli.conversation_history[0] + assert scaffold["display_kind"] == "hidden" + assert "REAL ASK" not in scaffold["content"] + assert scaffold["_db_persisted"] is True + active = db.get_messages_as_conversation(cli.session_id, include_row_ids=True) + assert len(active) == 1 + assert active[0]["content"] == scaffold["content"] + assert active[0]["_row_id"] == scaffold["_row_id"] + assert cli.agent._session_messages is cli.conversation_history + assert cli.agent._last_flushed_db_idx == 1 + assert cli.agent._db_flush_scan_prefix == cli.conversation_history + db.close() + + +def test_retry_last_rejects_media_before_db_or_memory_mutation(): + cli = _make_cli() + db = MagicMock() + cli._session_db = db + history = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "look again"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,AA"}}, + ], + }, + {"role": "assistant", "content": "old answer"}, + ] + cli.conversation_history = history + + 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() + + +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") + cli._session_db = db + history = [ + {"role": "user", "content": "retry me"}, + {"role": "assistant", "content": "old answer"}, + ] + cli.conversation_history = history + + assert cli.retry_last() is None + assert cli.conversation_history is history + + +def test_undo_last_prefills_live_text_and_retains_durable_scaffold(tmp_path): + cli = _make_cli() + cli._session_db.close() + db = SessionDB(db_path=tmp_path / "state.db") + cli._session_db = db + cli.session_id = "cli-carrier-undo" + db.create_session(cli.session_id, source="cli") + db.append_message(cli.session_id, "user", "older ask") + db.append_message(cli.session_id, "assistant", "older answer") + db.append_message(cli.session_id, "user", _composite_carrier()["content"]) + db.append_message(cli.session_id, "assistant", "failed answer") + cli.conversation_history = db.get_messages_as_conversation(cli.session_id) + cli._prefill_input_buffer = MagicMock() + cli.agent = SimpleNamespace( + _session_messages=cli.conversation_history, + _last_flushed_db_idx=len(cli.conversation_history), + _db_flush_scan_prefix=list(cli.conversation_history), + _invalidate_system_prompt=MagicMock(), + _memory_manager=None, + ) + + cli.undo_last() + + cli._prefill_input_buffer.assert_called_once_with("REAL ASK") + assert [m.get("content") for m in cli.conversation_history[:2]] == [ + "older ask", + "older answer", + ] + scaffold = cli.conversation_history[2] + assert scaffold["display_kind"] == "hidden" + assert "REAL ASK" not in scaffold["content"] + assert scaffold["_db_persisted"] is True + active = db.get_messages_as_conversation(cli.session_id, include_row_ids=True) + assert active[2]["_row_id"] == scaffold["_row_id"] + assert active[2]["content"] == scaffold["content"] + assert cli.agent._session_messages is cli.conversation_history + assert cli.agent._last_flushed_db_idx == 3 + db.close() diff --git a/tests/gateway/test_retry_replacement.py b/tests/gateway/test_retry_replacement.py index c83fafbe46..657c89d376 100644 --- a/tests/gateway/test_retry_replacement.py +++ b/tests/gateway/test_retry_replacement.py @@ -1,15 +1,31 @@ -"""Regression tests for /retry replacement semantics.""" +"""Regression tests for /retry replacement and carrier-aware undo semantics.""" +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import pytest +from agent.context_compressor import ( + HISTORICAL_TASK_HEADING, + SUMMARY_PREFIX, + _SUMMARY_END_MARKER, +) from gateway.config import GatewayConfig from gateway.platforms.base import MessageEvent, MessageType from gateway.run import GatewayRunner from gateway.session import SessionStore +def _composite_carrier(ask="REAL ASK"): + return { + "role": "user", + "content": ( + f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n" + f"{_SUMMARY_END_MARKER}\n\n{ask}" + ), + } + + @pytest.mark.asyncio async def test_gateway_retry_replaces_last_user_turn_in_transcript(tmp_path, monkeypatch): # Pin DEFAULT_DB_PATH so SessionDB() doesn't write to the real ~/.hermes/state.db. @@ -67,6 +83,191 @@ async def test_gateway_retry_replaces_last_user_turn_in_transcript(tmp_path, mon ] +@pytest.mark.asyncio +async def test_gateway_retry_redispatches_live_carrier_text_and_keeps_scaffold( + tmp_path, monkeypatch +): + import hermes_state + monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db") + + config = GatewayConfig() + store = SessionStore(sessions_dir=tmp_path, config=config) + session_id = "retry-carrier-session" + store._db.create_session(session_id=session_id, source="test") + store._db.append_message(session_id, "user", "older ask") + store._db.append_message(session_id, "assistant", "older answer") + store._db.append_message(session_id, "user", _composite_carrier()["content"]) + store._db.append_message(session_id, "assistant", "failed answer") + + gw = GatewayRunner.__new__(GatewayRunner) + gw.config = config + gw.session_store = store + session_entry = MagicMock(session_id=session_id, last_prompt_tokens=123) + gw.session_store.get_or_create_session = MagicMock(return_value=session_entry) + + async def fake_handle_message(event): + assert event.text == "REAL ASK" + active = store.load_transcript(session_id) + assert [m.get("content") for m in active[:2]] == ["older ask", "older answer"] + scaffold = active[2] + assert scaffold["display_kind"] == "hidden" + assert "REAL ASK" not in scaffold["content"] + return "new answer" + + gw._handle_message = AsyncMock(side_effect=fake_handle_message) + + result = await gw._handle_retry_command( + MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock()) + ) + + assert result == "new answer" + assert session_entry.last_prompt_tokens == 0 + gw._handle_message.assert_awaited_once() + archived = [ + row + for row in store._db.get_messages(session_id, include_inactive=True) + if not row["active"] + ] + assert [row["content"] for row in archived] == [ + _composite_carrier()["content"], + "failed answer", + ] + + +@pytest.mark.asyncio +async def test_gateway_retry_does_not_rewind_a_newer_plain_turn( + tmp_path, monkeypatch +): + """The carrier selected for retry must still be latest at commit time.""" + import hermes_state + monkeypatch.setattr(hermes_state, "DEFAULT_DB_PATH", tmp_path / "state.db") + + config = GatewayConfig() + store = SessionStore(sessions_dir=tmp_path, config=config) + session_id = "retry-carrier-race-session" + store._db.create_session(session_id=session_id, source="test") + store._db.append_message(session_id, "user", _composite_carrier()["content"]) + store._db.append_message(session_id, "assistant", "failed answer") + + gw = GatewayRunner.__new__(GatewayRunner) + gw.config = config + gw.session_store = store + session_entry = MagicMock(session_id=session_id, last_prompt_tokens=123) + gw.session_store.get_or_create_session = MagicMock(return_value=session_entry) + original_rewind = store.rewind_session + + def append_newer_turn_then_rewind(*args, **kwargs): + store._db.append_message(session_id, "user", "newer ask") + store._db.append_message(session_id, "assistant", "newer answer") + return original_rewind(*args, **kwargs) + + monkeypatch.setattr(store, "rewind_session", append_newer_turn_then_rewind) + gw._handle_message = AsyncMock() + + result = await gw._handle_retry_command( + MessageEvent(text="/retry", message_type=MessageType.TEXT, source=MagicMock()) + ) + + assert result.startswith("Retry failed;") + assert session_entry.last_prompt_tokens == 123 + gw._handle_message.assert_not_awaited() + assert [ + message.get("content") + for message in store.load_transcript(session_id) + if message.get("role") == "user" + ] == [_composite_carrier()["content"], "newer ask"] + + +@pytest.mark.asyncio +async def test_gateway_retry_rejects_media_before_redispatch_or_token_reset(): + 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=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "look again"}, + {"type": "image_url", "image_url": {"url": "image"}}, + ], + }, + {"role": "assistant", "content": "old answer"}, + ] + ), + rewrite_transcript=AsyncMock(return_value=True), + ) + 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.startswith("Cannot retry that message safely:") + assert session_entry.last_prompt_tokens == 123 + gw._handle_message.assert_not_awaited() + facade.rewrite_transcript.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_gateway_retry_stops_when_transcript_rewrite_fails(): + 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=[ + {"role": "user", "content": "retry me"}, + {"role": "assistant", "content": "old answer"}, + ] + ), + rewrite_transcript=AsyncMock(return_value=False), + ) + 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.startswith("Retry failed;") + assert session_entry.last_prompt_tokens == 123 + gw._handle_message.assert_not_awaited() + facade.rewrite_transcript.assert_awaited_once() + + +def test_gateway_undo_prefills_live_carrier_text_and_keeps_scaffold( + 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 = "undo-carrier-session" + store._db.create_session(session_id=session_id, source="test") + store._db.append_message(session_id, "user", _composite_carrier()["content"]) + store._db.append_message(session_id, "assistant", "failed answer") + + result = store.rewind_session(session_id) + + assert result["target_text"] == "REAL ASK" + assert result["rewound_count"] == 2 + active = store._db.get_messages_as_conversation( + session_id, include_row_ids=True + ) + assert len(active) == 1 + assert active[0]["display_kind"] == "hidden" + assert "REAL ASK" not in active[0]["content"] + + @pytest.mark.asyncio async def test_gateway_retry_preserves_archived_compaction_rows_when_probe_fails( tmp_path, monkeypatch diff --git a/tests/gateway/test_undo_rewind_session.py b/tests/gateway/test_undo_rewind_session.py index b95f34622a..10df30a431 100644 --- a/tests/gateway/test_undo_rewind_session.py +++ b/tests/gateway/test_undo_rewind_session.py @@ -54,3 +54,32 @@ def test_rewind_n_turns(store): assert len(store.load_transcript(sid)) == 2 # q1,a1 +def test_rewind_fails_closed_when_transcript_changes_after_snapshot( + store, monkeypatch +): + sid = _seed(store, "gw-cas", turns=2) + sibling = SessionDB(db_path=store._db.db_path) + original_rewind = store._db.rewind_to_message + + def _append_then_rewind(*args, **kwargs): + sibling.append_message(sid, "assistant", "concurrent tail") + return original_rewind(*args, **kwargs) + + monkeypatch.setattr(store._db, "rewind_to_message", _append_then_rewind) + + 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), + ("concurrent tail", 1), + ] + sibling.close() + diff --git a/tests/hermes_state/test_composite_carrier_rewind.py b/tests/hermes_state/test_composite_carrier_rewind.py new file mode 100644 index 0000000000..444f58495f --- /dev/null +++ b/tests/hermes_state/test_composite_carrier_rewind.py @@ -0,0 +1,238 @@ +"""Transactional persistence contracts for composite compaction carriers.""" + +from __future__ import annotations + +import pytest + +from agent.context_compressor import ( + HISTORICAL_TASK_HEADING, + SUMMARY_PREFIX, + _MERGED_SUMMARY_DELIMITER, + _SUMMARY_END_MARKER, +) +from hermes_state import ( + CompressionSessionClosedError, + SessionCompressionInProgressError, + SessionDB, +) + + +def _carrier(ask: str = "REAL ASK") -> str: + return ( + f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n" + f"{_SUMMARY_END_MARKER}\n\n{ask}" + ) + + +@pytest.fixture() +def db(tmp_path): + state = SessionDB(db_path=tmp_path / "state.db") + yield state + state.close() + + +def _session_counts(db: SessionDB, session_id: str) -> tuple[int, int, int]: + row = db._conn.execute( + "SELECT message_count, tool_call_count, rewind_count " + "FROM sessions WHERE id = ?", + (session_id,), + ).fetchone() + return row["message_count"], row["tool_call_count"], row["rewind_count"] + + +def _row_state(db: SessionDB, session_id: str) -> list[tuple]: + return [ + tuple(row) + for row in db._conn.execute( + "SELECT id, role, content, active, display_kind " + "FROM messages WHERE session_id = ? ORDER BY id", + (session_id,), + ).fetchall() + ] + + +def _active_ids(db: SessionDB, session_id: str) -> list[int]: + return [ + int(message["_row_id"]) + for message in db.get_messages_as_conversation( + session_id, include_row_ids=True + ) + ] + + +def test_composite_rewind_archives_tail_and_inserts_its_hidden_scaffold(db): + sid = "carrier-rewind" + db.create_session(sid, source="tui") + db.append_message(sid, "user", "older ask") + db.append_message( + sid, + "assistant", + None, + tool_calls=[{"id": "call-1", "function": {"name": "terminal"}}], + ) + db.append_message(sid, "tool", "ok", tool_call_id="call-1") + target_id = db.append_message(sid, "user", _carrier()) + db.append_message(sid, "assistant", "failed") + expected_active_ids = _active_ids(db, sid) + + result = db.rewind_to_message( + sid, + target_id, + preserve_compaction_handoff=True, + expected_active_ids=expected_active_ids, + expected_target_content="REAL ASK", + ) + + assert result["rewound_count"] == 2 + assert result["replacement_message_id"] == result["new_head_id"] + active = db.get_messages_as_conversation(sid, include_row_ids=True) + assert len(active) == 4 + assert active[-1]["_row_id"] == result["replacement_message_id"] + assert active[-1]["display_kind"] == "hidden" + assert SUMMARY_PREFIX in active[-1]["content"] + assert "REAL ASK" not in active[-1]["content"] + archived = db._conn.execute( + "SELECT active FROM messages WHERE id IN (?, ?) ORDER BY id", + (target_id, target_id + 1), + ).fetchall() + assert [row[0] for row in archived] == [0, 0] + assert _session_counts(db, sid) == (4, 1, 1) + + +def test_default_rewind_return_shape_and_active_counters_remain_compatible(db): + sid = "default-rewind" + db.create_session(sid, source="cli") + db.append_message(sid, "user", "first") + db.append_message(sid, "assistant", "answer") + target_id = db.append_message(sid, "user", "second") + db.append_message( + sid, + "assistant", + None, + tool_calls=[{"id": "call-2", "function": {"name": "terminal"}}], + ) + + result = db.rewind_to_message(sid, target_id) + + assert set(result) == {"rewound_count", "target_message", "new_head_id"} + assert result["rewound_count"] == 2 + assert _session_counts(db, sid) == (2, 0, 1) + + +def test_guarded_composite_rewind_rejects_append_without_inserting_scaffold(db): + sid = "guarded-rewind-append" + db.create_session(sid, source="cli") + db.append_message(sid, "user", "first") + db.append_message(sid, "assistant", "answer") + target_id = db.append_message(sid, "user", _carrier()) + db.append_message(sid, "assistant", "failed") + snapshot = db.get_messages_as_conversation(sid, include_row_ids=True) + expected_active_ids = [int(message["_row_id"]) for message in snapshot] + assert snapshot[-2]["_row_id"] == target_id + + # Deterministic validation -> write race: a sibling writer commits after + # the snapshot but before rewind_to_message begins its write transaction. + sibling = SessionDB(db_path=db.db_path) + sibling.append_message(sid, "assistant", "concurrent append") + sibling.close() + before_rows = _row_state(db, sid) + before_counts = _session_counts(db, sid) + + with pytest.raises(RuntimeError, match="active transcript changed"): + db.rewind_to_message( + sid, + target_id, + preserve_compaction_handoff=True, + expected_active_ids=expected_active_ids, + expected_target_content=_carrier(), + ) + + assert _row_state(db, sid) == before_rows + assert _session_counts(db, sid) == before_counts + +def test_guarded_rewind_rejects_selected_target_content_change(db): + sid = "guarded-rewind-in-place" + db.create_session(sid, source="cli") + db.append_message(sid, "user", "first") + db.append_message(sid, "assistant", "answer") + target_id = db.append_message(sid, "user", "second") + db.append_message(sid, "assistant", "failed") + expected_active_ids = _active_ids(db, sid) + + sibling = SessionDB(db_path=db.db_path) + sibling._execute_write( + lambda conn: conn.execute( + "UPDATE messages SET content = ? WHERE id = ?", + ("changed second", target_id), + ) + ) + sibling.close() + before_rows = _row_state(db, sid) + before_counts = _session_counts(db, sid) + + with pytest.raises(RuntimeError, match="rewind target changed"): + db.rewind_to_message( + sid, + target_id, + expected_active_ids=expected_active_ids, + expected_target_content="second", + ) + + assert _row_state(db, sid) == before_rows + assert _session_counts(db, sid) == before_counts + + +def test_guarded_rewind_ignores_reaction_metadata_change(db): + sid = "guarded-rewind-reaction" + db.create_session(sid, source="cli") + target_id = db.append_message(sid, "user", "second") + db.append_message(sid, "assistant", "failed") + expected_active_ids = _active_ids(db, sid) + + assert db.set_message_reaction(sid, target_id + 1, "👍", author="user") + + result = db.rewind_to_message( + sid, + target_id, + expected_active_ids=expected_active_ids, + expected_target_content="second", + ) + + assert result["rewound_count"] == 2 + assert db.get_messages_as_conversation(sid) == [] + + +def test_rewind_guard_rejects_foreign_live_compression_without_any_change(db): + sid = "locked-rewind" + db.create_session(sid, source="tui") + target_id = db.append_message(sid, "user", _carrier()) + db.append_message(sid, "assistant", "failed") + assert db.try_acquire_compression_lock(sid, "foreign-writer", ttl_seconds=60) + before_rows = _row_state(db, sid) + before_counts = _session_counts(db, sid) + + with pytest.raises(SessionCompressionInProgressError): + db.rewind_to_message( + sid, target_id, preserve_compaction_handoff=True + ) + + assert _row_state(db, sid) == before_rows + assert _session_counts(db, sid) == before_counts + + +def test_rewind_guard_rejects_compression_ended_parent_without_any_change(db): + sid = "closed-rewind" + db.create_session(sid, source="tui") + target_id = db.append_message(sid, "user", _carrier()) + db.append_message(sid, "assistant", "failed") + db.end_session(sid, "compression") + before_rows = _row_state(db, sid) + before_counts = _session_counts(db, sid) + + with pytest.raises(CompressionSessionClosedError): + db.rewind_to_message( + sid, target_id, preserve_compaction_handoff=True + ) + + assert _row_state(db, sid) == before_rows + assert _session_counts(db, sid) == before_counts diff --git a/tests/run_agent/test_identity_flush.py b/tests/run_agent/test_identity_flush.py index 5c8a699ead..8cb1e54d6d 100644 --- a/tests/run_agent/test_identity_flush.py +++ b/tests/run_agent/test_identity_flush.py @@ -31,6 +31,55 @@ def _contents(db, session_id=SESSION_ID): class TestIdentityFlush: + def test_summary_flush_hides_pure_handoff_but_not_composite_live_ask(self): + from agent.context_compressor import ( + COMPRESSED_SUMMARY_METADATA_KEY, + HISTORICAL_TASK_HEADING, + SUMMARY_PREFIX, + _SUMMARY_END_MARKER, + ) + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmpdir: + db = SessionDB(db_path=Path(tmpdir) / "t.db") + try: + agent = _make_agent(db) + scaffold = ( + f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold\n\n" + f"{_SUMMARY_END_MARKER}" + ) + messages = [ + { + "role": "user", + "content": scaffold, + COMPRESSED_SUMMARY_METADATA_KEY: True, + }, + { + "role": "user", + "content": scaffold + "\n\nREAL ASK", + COMPRESSED_SUMMARY_METADATA_KEY: True, + }, + { + "role": "assistant", + "content": scaffold, + COMPRESSED_SUMMARY_METADATA_KEY: True, + }, + ] + + agent._flush_messages_to_session_db(messages, []) + + rows = db._conn.execute( + "SELECT content, display_kind FROM messages " + "WHERE session_id = ? ORDER BY id", + (SESSION_ID,), + ).fetchall() + assert rows[0]["display_kind"] == "hidden" + assert rows[1]["display_kind"] is None + assert rows[1]["content"].endswith("REAL ASK") + assert rows[2]["display_kind"] == "hidden" + finally: + db.close() + def test_repair_shrunk_messages_below_history_length_still_persists_assistant(self): """When repair shortens messages below conversation_history, don't slice empty.""" from hermes_state import SessionDB diff --git a/tests/run_agent/test_thinking_only_sanitizer.py b/tests/run_agent/test_thinking_only_sanitizer.py index d6875084ea..f6190fab72 100644 --- a/tests/run_agent/test_thinking_only_sanitizer.py +++ b/tests/run_agent/test_thinking_only_sanitizer.py @@ -93,6 +93,17 @@ class TestDropThinkingOnlyAndMergeUsers: # Should return the original list untouched (identity) when no changes. assert out is msgs + def test_adjacent_users_merge_even_when_no_thinking_row_was_dropped(self): + scaffold = {"role": "user", "content": "SUMMARY SCAFFOLD"} + live_ask = {"role": "user", "content": "REAL ASK"} + msgs = [scaffold, live_ask] + + out = AIAgent._drop_thinking_only_and_merge_users(msgs) + + assert out == [{"role": "user", "content": "SUMMARY SCAFFOLD\n\nREAL ASK"}] + assert scaffold["content"] == "SUMMARY SCAFFOLD" + assert live_ask["content"] == "REAL ASK" + def test_preserves_alternation_after_drop(self): msgs = [ diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index 1407c9586f..384abd04a0 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -9140,6 +9140,7 @@ def test_rollback_restore_truncates_from_real_user_turn_not_marker(monkeypatch): server._sessions["sid"] = _session( agent=types.SimpleNamespace(_checkpoint_mgr=_Mgr()), history=list(history), + session_key="", ) try: resp = server.handle_request( @@ -9200,6 +9201,7 @@ def test_rollback_restore_skips_legacy_compaction_handoff(monkeypatch): server._sessions["sid"] = _session( agent=types.SimpleNamespace(_checkpoint_mgr=_Mgr()), history=list(history), + session_key="", ) try: resp = server.handle_request( @@ -9222,6 +9224,85 @@ def test_rollback_restore_skips_legacy_compaction_handoff(monkeypatch): # ── session.steer ──────────────────────────────────────────────────── +def test_rollback_restore_preserves_composite_carrier_scaffold(monkeypatch, tmp_path): + """A checkpoint restore drops the live ask but keeps compacted context.""" + from agent.context_compressor import ( + HISTORICAL_TASK_HEADING, + SUMMARY_PREFIX, + _SUMMARY_END_MARKER, + ) + from hermes_state import SessionDB + + class _Mgr: + enabled = True + + def list_checkpoints(self, cwd): + return [{"hash": "abc123"}] + + def restore(self, cwd, target, file_path=None): + return {"success": True, "message": "restored"} + + carrier = { + "role": "user", + "content": ( + f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n" + f"{_SUMMARY_END_MARKER}\n\nREAL ASK" + ), + } + db = SessionDB(db_path=tmp_path / "state.db") + db.create_session("rollback-carrier", source="tui") + db.append_message("rollback-carrier", "user", carrier["content"]) + db.append_message("rollback-carrier", "assistant", "answer") + durable = db.get_messages_as_conversation("rollback-carrier") + agent = types.SimpleNamespace( + _checkpoint_mgr=_Mgr(), + _session_messages=list(durable), + _last_flushed_db_idx=len(durable), + _db_flush_scan_prefix=list(durable), + ) + server._sessions["sid"] = _session( + agent=agent, + history=list(durable), + session_key="rollback-carrier", + ) + monkeypatch.setattr(server, "_get_db", lambda: db) + try: + resp = server.handle_request( + { + "id": "1", + "method": "rollback.restore", + "params": {"session_id": "sid", "hash": "abc123"}, + } + ) + + assert "result" in resp, resp + assert resp["result"]["success"] is True + assert resp["result"]["history_removed"] == 2 + remaining = server._sessions["sid"]["history"] + assert len(remaining) == 1 + assert remaining[0]["display_kind"] == "hidden" + assert SUMMARY_PREFIX in remaining[0]["content"] + assert "REAL ASK" not in remaining[0]["content"] + cold = db.get_messages_as_conversation( + "rollback-carrier", include_row_ids=True + ) + assert len(cold) == 1 + assert cold[0]["content"] == remaining[0]["content"] + assert cold[0]["display_kind"] == "hidden" + assert cold[0]["_row_id"] == remaining[0]["_row_id"] + assert agent._session_messages == remaining + assert agent._last_flushed_db_idx == 1 + assert agent._db_flush_scan_prefix == remaining + inactive = db.get_messages_as_conversation( + "rollback-carrier", include_inactive=True + ) + assert any("REAL ASK" in str(message.get("content")) for message in inactive) + assert any(message.get("content") == "answer" for message in inactive) + finally: + server._sessions.pop("sid", None) + db.close() + + def test_session_steer_calls_agent_steer_when_agent_supports_it(): """The TUI RPC method must call agent.steer(text) and return a queued status without touching interrupt state. @@ -9502,6 +9583,7 @@ def test_session_undo_allowed_when_idle(): """Regression guard: when not running, /undo still works.""" server._sessions["sid"] = _session( running=False, + session_key="", history=[ {"role": "user", "content": "hi"}, {"role": "assistant", "content": "hello"}, diff --git a/tests/tui_gateway/test_composite_carrier_rewind.py b/tests/tui_gateway/test_composite_carrier_rewind.py new file mode 100644 index 0000000000..6af08745f8 --- /dev/null +++ b/tests/tui_gateway/test_composite_carrier_rewind.py @@ -0,0 +1,587 @@ +"""Regression contracts for TUI rewind of live compaction carriers.""" + +from __future__ import annotations + +import threading +from types import SimpleNamespace + +import pytest + +from agent.context_compressor import ( + HISTORICAL_TASK_HEADING, + SUMMARY_PREFIX, + _SUMMARY_END_MARKER, +) +from hermes_state import SessionDB +from tui_gateway import server + + +def _composite_carrier() -> dict: + return { + "role": "user", + "content": ( + f"{SUMMARY_PREFIX}\n{HISTORICAL_TASK_HEADING}\nold task\n\n" + f"{_SUMMARY_END_MARKER}\n\nREAL ASK" + ), + } + + +@pytest.fixture() +def carrier_session(tmp_path): + old_db = server._db + db = SessionDB(db_path=tmp_path / "state.db") + installed_ids: list[str] = [] + + def install(history: list[dict]): + sid = f"carrier-sid-{len(installed_ids)}" + session_key = f"carrier-session-{len(installed_ids)}" + installed_ids.append(sid) + db.create_session(session_key, source="tui") + for message in history: + db.append_message( + session_key, + message["role"], + message.get("content"), + ) + durable = db.get_messages_as_conversation(session_key) + agent = SimpleNamespace( + _session_messages=list(durable), + _last_flushed_db_idx=len(durable), + _db_flush_scan_prefix=list(durable), + ) + session = { + "agent": agent, + "attached_images": [], + "history": list(durable), + "history_lock": threading.Lock(), + "history_version": 0, + "running": False, + "session_key": session_key, + } + server._sessions[sid] = session + return sid, session_key, session + + server._db = db + yield db, install + for sid in installed_ids: + server._sessions.pop(sid, None) + server._db = old_db + db.close() + + +def _dispatch(sid: str, name: str) -> dict: + return server._methods["command.dispatch"]( + "request-id", + {"session_id": sid, "name": name, "arg": ""}, + ) + + +def _session_undo(sid: str) -> dict: + return server._methods["session.undo"]( + "request-id", + {"session_id": sid}, + ) + + +def _assert_scaffold_preserved( + db: SessionDB, + session_key: str, + session: dict, + *, + prefix_len: int = 0, +) -> None: + active = db.get_messages_as_conversation(session_key, include_row_ids=True) + assert len(active) == prefix_len + 1 + scaffold = active[prefix_len] + assert scaffold["role"] == "user" + assert scaffold["display_kind"] == "hidden" + assert SUMMARY_PREFIX in scaffold["content"] + assert "REAL ASK" not in scaffold["content"] + assert session["history"][prefix_len]["content"] == scaffold["content"] + assert session["history"][prefix_len]["display_kind"] == "hidden" + + +def test_retry_selects_the_live_ask_inside_a_force_user_leading_carrier( + carrier_session, +): + db, install = carrier_session + sid, session_key, session = install( + [_composite_carrier(), {"role": "assistant", "content": "failed"}] + ) + + response = _dispatch(sid, "retry") + + assert response["result"] == {"type": "send", "message": "REAL ASK"} + _assert_scaffold_preserved(db, session_key, session) + + +@pytest.mark.parametrize("command", ["retry", "undo"]) +def test_rewind_matches_cold_sanitized_carrier_to_unchanged_warm_ask( + carrier_session, command +): + db, install = carrier_session + carrier = _composite_carrier() + carrier["content"] = carrier["content"].replace( + "REAL ASK", + " REAL ASK\n\n\nprivate\n ", + ) + sid, session_key, session = install( + [carrier, {"role": "assistant", "content": "failed"}] + ) + db._conn.execute( + "UPDATE messages SET api_content = ? " + "WHERE session_id = ? AND role = 'user'", + (carrier["content"], session_key), + ) + db._conn.commit() + # A live agent still has the raw wire form while a cold DB projection has + # already applied the role-aware sanitize_context(...).strip() rule and + # retained the raw provider wire in api_content. + session["history"][0] = carrier.copy() + session["agent"]._session_messages = list(session["history"]) + + response = _dispatch(sid, command) + + assert response["result"]["message"] == "REAL ASK" + _assert_scaffold_preserved(db, session_key, session) + + +def test_retry_fails_closed_when_transcript_changes_after_snapshot( + carrier_session, monkeypatch +): + db, install = carrier_session + sid, session_key, session = install( + [_composite_carrier(), {"role": "assistant", "content": "failed"}] + ) + sibling = SessionDB(db_path=db.db_path) + original_rewind = db.rewind_to_message + + def _append_then_rewind(*args, **kwargs): + sibling.append_message(session_key, "assistant", "concurrent tail") + return original_rewind(*args, **kwargs) + + monkeypatch.setattr(db, "rewind_to_message", _append_then_rewind) + before_history = [dict(message) for message in session["history"]] + + response = _dispatch(sid, "retry") + + assert response["error"]["code"] == 5008 + assert "active transcript changed" in response["error"]["message"] + assert session["history"] == before_history + rows = db._conn.execute( + "SELECT content, active, display_kind FROM messages " + "WHERE session_id = ? ORDER BY id", + (session_key,), + ).fetchall() + assert [tuple(row) for row in rows] == [ + (_composite_carrier()["content"], 1, None), + ("failed", 1, None), + ("concurrent tail", 1, None), + ] + sibling.close() + + +@pytest.mark.parametrize("command", ["retry", "undo"]) +def test_rewind_allows_database_only_reaction_metadata_change( + carrier_session, command +): + db, install = carrier_session + sid, session_key, session = install( + [ + {"role": "user", "content": "OLDER ASK"}, + {"role": "assistant", "content": "older answer"}, + _composite_carrier(), + {"role": "assistant", "content": "failed"}, + ] + ) + older_answer = next( + row + for row in db.get_messages(session_key) + if row["role"] == "assistant" and row["content"] == "older answer" + ) + assert db.set_message_reaction( + session_key, older_answer["id"], "👍", author="user" + ) + + response = _dispatch(sid, command) + + assert response["result"]["message"] == "REAL ASK" + _assert_scaffold_preserved(db, session_key, session, prefix_len=2) + + +def test_retry_ignores_buried_ephemeral_scaffolding_missing_from_db( + carrier_session, +): + db, install = carrier_session + sid, session_key, session = install( + [ + {"role": "user", "content": "OLDER ASK"}, + {"role": "assistant", "content": "older answer"}, + _composite_carrier(), + {"role": "assistant", "content": "failed"}, + ] + ) + session["history"].insert( + 2, + { + "role": "user", + "content": "internal recovery nudge", + "_dropped_toolcall_nudge": True, + }, + ) + session["agent"]._session_messages = list(session["history"]) + + response = _dispatch(sid, "retry") + + assert response["result"] == {"type": "send", "message": "REAL ASK"} + _assert_scaffold_preserved(db, session_key, session, prefix_len=2) + + +def test_retry_drops_buried_ephemeral_scaffolding_from_the_warm_prefix( + carrier_session, +): + db, install = carrier_session + sid, session_key, session = install( + [ + {"role": "user", "content": "OLDER ASK"}, + {"role": "assistant", "content": "candidate answer"}, + {"role": "assistant", "content": "verified answer"}, + _composite_carrier(), + {"role": "assistant", "content": "failed"}, + ] + ) + session["history"].insert( + 2, + { + "role": "user", + "content": "[System: verify before stopping]", + "_verification_stop_synthetic": True, + }, + ) + session["agent"]._session_messages = list(session["history"]) + + response = _dispatch(sid, "retry") + + assert response["result"] == {"type": "send", "message": "REAL ASK"} + assert [message.get("content") for message in session["history"][:3]] == [ + "OLDER ASK", + "candidate answer", + "verified answer", + ] + active = db.get_messages_as_conversation(session_key, include_row_ids=True) + # Alternation repair is a model/memory projection; the two durable source + # rows remain independently recoverable ahead of the inserted scaffold. + assert [message.get("content") for message in active[:3]] == [ + "OLDER ASK", + "candidate answer", + "verified answer", + ] + assert active[3]["display_kind"] == "hidden" + assert "REAL ASK" not in active[3]["content"] + assert session["history"][3]["display_kind"] == "hidden" + assert "REAL ASK" not in session["history"][3]["content"] + + +def test_retry_preserves_older_warm_media_while_targeting_plain_ask( + carrier_session, +): + db, install = carrier_session + sid, session_key, session = install( + [ + {"role": "user", "content": "look\n[screenshot]"}, + {"role": "assistant", "content": "seen"}, + _composite_carrier(), + {"role": "assistant", "content": "failed"}, + ] + ) + session["history"][0]["content"] = [ + {"type": "text", "text": "look"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,x"}}, + ] + session["agent"]._session_messages = list(session["history"]) + + response = _dispatch(sid, "retry") + + assert response["result"] == {"type": "send", "message": "REAL ASK"} + assert isinstance(session["history"][0]["content"], list) + _assert_scaffold_preserved(db, session_key, session, prefix_len=2) + + +def test_undo_targets_the_composite_carrier_not_an_older_user_turn( + carrier_session, +): + db, install = carrier_session + sid, session_key, session = install( + [ + {"role": "user", "content": "OLDER ASK"}, + {"role": "assistant", "content": "older answer"}, + _composite_carrier(), + {"role": "assistant", "content": "failed"}, + ] + ) + + response = _dispatch(sid, "undo") + + assert response["result"]["type"] == "prefill" + assert response["result"]["message"] == "REAL ASK" + active = db.get_messages_as_conversation(session_key, include_row_ids=True) + assert [message.get("content") for message in active[:2]] == [ + "OLDER ASK", + "older answer", + ] + _assert_scaffold_preserved(db, session_key, session, prefix_len=2) + + +def test_undo_rewinds_media_placeholder_without_treating_it_as_retry( + carrier_session, +): + db, install = carrier_session + carrier = _composite_carrier() + carrier["content"] = carrier["content"].replace( + "REAL ASK", "look\n[screenshot]" + ) + sid, session_key, session = install( + [carrier, {"role": "assistant", "content": "seen"}] + ) + + response = _dispatch(sid, "undo") + + assert response["result"]["type"] == "prefill" + assert response["result"]["message"] == "look\n[screenshot]" + active = db.get_messages_as_conversation(session_key) + assert len(active) == 1 + assert active[0]["display_kind"] == "hidden" + assert "look\n[screenshot]" not in active[0]["content"] + assert len(session["history"]) == 1 + assert session["history"][0]["content"] == active[0]["content"] + assert session["history"][0]["display_kind"] == "hidden" + + +def test_session_undo_preserves_the_composite_carriers_scaffold(carrier_session): + db, install = carrier_session + sid, session_key, session = install( + [_composite_carrier(), {"role": "assistant", "content": "answer"}] + ) + + response = _session_undo(sid) + + assert response["result"]["removed"] == 2 + _assert_scaffold_preserved(db, session_key, session) + + +def test_history_projection_unwraps_composite_and_hides_sole_handoff(): + composite = {**_composite_carrier(), "_row_id": 7} + sole_handoff = { + **composite, + "content": composite["content"].split("\n\nREAL ASK", 1)[0], + } + + assert server._history_to_messages( + [composite, {"role": "user", "content": "newer ask", "_row_id": 9}] + ) == [ + {"role": "user", "text": "REAL ASK", "row_id": 7}, + {"role": "user", "text": "newer ask", "row_id": 9}, + ] + assert server._history_to_messages([sole_handoff]) == [] + + +def test_retry_preserves_literal_media_like_text(carrier_session): + db, install = carrier_session + carrier = _composite_carrier() + carrier["content"] = carrier["content"].replace( + "REAL ASK", "inspect [image|ybres:RID]" + ) + sid, session_key, session = install( + [carrier, {"role": "assistant", "content": "failed"}] + ) + response = _dispatch(sid, "retry") + + assert response["result"] == { + "type": "send", + "message": "inspect [image|ybres:RID]", + } + _assert_scaffold_preserved(db, session_key, session) + + +def test_retry_rejects_pending_attachments_before_mutating_history(carrier_session): + db, install = carrier_session + sid, session_key, session = install( + [_composite_carrier(), {"role": "assistant", "content": "failed"}] + ) + session["attached_images"] = ["/tmp/pending.png"] + before_memory = list(session["history"]) + + response = _dispatch(sid, "retry") + + assert response["error"]["code"] == 4018 + assert session["history"] == before_memory + assert len(db.get_messages_as_conversation(session_key)) == 2 + +def test_prompt_ordinal_rewind_preserves_scaffold_before_regeneration( + carrier_session, monkeypatch +): + db, install = carrier_session + sid, session_key, session = install( + [_composite_carrier(), {"role": "assistant", "content": "failed"}] + ) + seen = {} + + class _Agent: + _session_messages = list(session["history"]) + _last_flushed_db_idx = len(_session_messages) + _db_flush_scan_prefix = list(_session_messages) + + def run_conversation( + self, prompt, conversation_history=None, stream_callback=None, **_kwargs + ): + seen["prompt"] = prompt + seen["history"] = list(conversation_history or []) + return { + "final_response": "regenerated", + "messages": [ + *(conversation_history or []), + {"role": "user", "content": prompt}, + {"role": "assistant", "content": "regenerated"}, + ], + } + + class _ImmediateThread: + def __init__(self, target=None, daemon=None): + self._target = target + + def start(self): + self._target() + + session["agent"] = _Agent() + monkeypatch.setattr(server.threading, "Thread", _ImmediateThread) + monkeypatch.setattr(server, "_get_usage", lambda _agent: {}) + monkeypatch.setattr(server, "render_message", lambda *_args: "") + monkeypatch.setattr(server, "_emit", lambda *_args: None) + + response = server._methods["prompt.submit"]( + "request-id", + { + "session_id": sid, + "text": "EDITED ASK", + "truncate_before_user_ordinal": 0, + "confirm_truncate": True, + }, + ) + + assert response["result"]["status"] == "streaming" + assert seen["prompt"] == "EDITED ASK" + assert len(seen["history"]) == 1 + assert seen["history"][0]["display_kind"] == "hidden" + assert "REAL ASK" not in seen["history"][0]["content"] + active = db.get_messages_as_conversation(session_key, include_row_ids=True) + assert len(active) == 1 + assert active[0]["display_kind"] == "hidden" + + +def test_prompt_row_id_rewind_uses_profile_db_and_rebinds_survivors( + monkeypatch, tmp_path +): + profile_home = tmp_path / "profiles" / "work" + profile_home.mkdir(parents=True) + profile_db = SessionDB(db_path=profile_home / "state.db") + launch_db = SessionDB(db_path=tmp_path / "launch.db") + session_key = "profile-carrier-row-id" + profile_db.create_session(session_key, source="tui") + launch_db.create_session(session_key, source="tui") + launch_db.append_message(session_key, "user", "launch profile must stay untouched") + + carrier = _composite_carrier() + carrier["content"] = carrier["content"].replace( + "REAL ASK", + " REAL ASK\n\nprivate ", + ) + persisted = [ + {"role": "user", "content": "OLDER ASK"}, + {"role": "assistant", "content": "older answer"}, + carrier, + {"role": "assistant", "content": "failed"}, + ] + for message in persisted: + profile_db.append_message( + session_key, + message["role"], + message["content"], + ) + target_row_id = profile_db.get_messages_as_conversation( + session_key, include_row_ids=True + )[2]["_row_id"] + + sid = "profile-carrier-row-id-sid" + session = { + "agent": SimpleNamespace(), + "attached_images": [], + "history": [dict(message) for message in persisted], + "history_lock": threading.Lock(), + "history_version": 0, + "image_counter": 0, + "profile_home": str(profile_home), + "running": False, + "session_key": session_key, + "show_reasoning": False, + "slash_worker": None, + "tool_progress_mode": "all", + } + + class _DormantThread: + def __init__(self, target=None, daemon=None): + self._target = target + + def start(self): + pass + + old_db = server._db + server._db = launch_db + server._sessions[sid] = session + monkeypatch.setattr(server, "_start_agent_build", lambda *_args, **_kwargs: None) + monkeypatch.setattr(server, "_start_inflight_turn", lambda *_args, **_kwargs: None) + monkeypatch.setattr(server.threading, "Thread", _DormantThread) + + try: + response = server._methods["prompt.submit"]( + "request-id", + { + "session_id": sid, + "text": "EDITED ASK", + "truncate_before_row_id": target_row_id, + "truncate_before_user_ordinal": 1, + "confirm_truncate": True, + }, + ) + + assert response["result"]["status"] == "streaming" + survivor_ids = response["result"]["survivor_user_row_ids"] + assert len(survivor_ids) == 1 + assert isinstance(survivor_ids[0], int) + + active = profile_db.get_messages_as_conversation( + session_key, include_row_ids=True + ) + assert [message["content"] for message in active[:2]] == [ + "OLDER ASK", + "older answer", + ] + assert active[0]["_row_id"] == survivor_ids[0] + assert active[2]["display_kind"] == "hidden" + assert "REAL ASK" not in active[2]["content"] + assert "private" not in active[2]["content"] + assert [ + (message["role"], message["content"], message.get("display_kind")) + for message in session["history"] + ] == [ + (message["role"], message["content"], message.get("display_kind")) + for message in active + ] + assert [ + message["content"] + for message in launch_db.get_messages_as_conversation(session_key) + ] == ["launch profile must stay untouched"] + finally: + server._sessions.pop(sid, None) + server._db = old_db + profile_db.close() + launch_db.close() diff --git a/tui_gateway/methods_prompt.py b/tui_gateway/methods_prompt.py index 0e13b4cb24..ee8d93a6b0 100644 --- a/tui_gateway/methods_prompt.py +++ b/tui_gateway/methods_prompt.py @@ -14,11 +14,13 @@ _profile_scoped = _registry.profile_scoped def _history_user_indices(history: list) -> list: - """Indices of model-visible user turns (excludes display_kind timeline markers).""" + """Indices of canonical live-user turns, including composite carriers.""" + from agent.context_compressor import user_originated_turn_view + return [ i for i, m in enumerate(history) - if m.get("role") == "user" and not m.get("display_kind") + if user_originated_turn_view(m) is not None ] @@ -50,17 +52,28 @@ def _mem_db_pair_agrees(mem, db_msg) -> bool: return False if mem.get("role") != db_msg.get("role"): return False + if mem.get("role") == "user": + from agent.context_compressor import user_originated_turn_view + from agent.memory_manager import sanitize_context + + mem_view = user_originated_turn_view(mem) + db_view = user_originated_turn_view(db_msg) + if (mem_view is None) != (db_view is None): + return False + if mem_view is None: + return bool(mem.get("display_kind")) == bool( + db_msg.get("display_kind") + ) + mem_content = mem_view.get("content") + db_content = db_view.get("content") + if isinstance(mem_content, str) and isinstance(db_content, str): + if sanitize_context(mem_content).strip() != sanitize_context( + db_content + ).strip(): + return False + return True if bool(mem.get("display_kind")) != bool(db_msg.get("display_kind")): return False - if mem.get("role") == "user" and not mem.get("display_kind"): - mem_content = mem.get("content") - db_content = db_msg.get("content") - if ( - isinstance(mem_content, str) - and isinstance(db_content, str) - and mem_content.strip() != db_content.strip() - ): - return False return True @@ -90,20 +103,15 @@ def _resolve_truncate_row_id(session: dict, history: list, target_row_id: int): return None try: - db = _get_db() - except Exception: - db = None - if db is None: - return None - - get_conv = getattr(db, "get_messages_as_conversation", None) - if not callable(get_conv): - return None - - try: - db_history = get_conv( - session_key, repair_alternation=True, include_row_ids=True - ) + with _session_db(session) as db: + if db is None: + return None + get_conv = getattr(db, "get_messages_as_conversation", None) + if not callable(get_conv): + return None + db_history = get_conv( + session_key, repair_alternation=True, include_row_ids=True + ) except Exception: logger.debug( "prompt.submit: failed loading DB history for row_id %s session %s", @@ -379,7 +387,9 @@ def _(rid, params: dict) -> dict: or truncate_message_id is not None or truncate_row_id is not None ): - history = session.get("history", []) + history = _history_without_ephemeral_scaffolding( + session.get("history", []) + ) # Malformed params refuse first (4004), regardless of consent — # the historical ordinal-path precedence. @@ -504,7 +514,11 @@ def _(rid, params: dict) -> dict: # via replace_messages — an unrecoverable overwrite of the session DB. if ordinal < 0 or ordinal >= len(user_indices): return _err(rid, 4018, "target user message is no longer in session history") - truncated = history[: user_indices[ordinal]] + from agent.context_compressor import history_before_user_originated_turn + + truncated, _live_view = history_before_user_originated_turn( + history, user_indices[ordinal] + ) # Second gate, on top of confirm_truncate: ordinal 0 resolves to # history[:0] == [] and replace_messages() DELETEs every durable # row. A confirmed rewind that happens to erase the whole @@ -547,58 +561,39 @@ def _(rid, params: dict) -> dict: # new exchange is appended on top of the "undone" turns — durable # zombie history on resume, and the edit/regenerate never sticks. # Fail closed: refuse the turn and leave memory/DB unchanged. - if (db := _get_db()) is not None: - try: - # active_only=True: replace only the live (active=1) rows. - # In-place compaction (#38763) keeps the pre-compaction - # transcript as active=0/compacted=1 rows under this same - # session key; a bare replace_messages() would DELETE that - # durable archive on every edit/regenerate — the same bug - # class #80216 fixed for /retry. On an uncompacted session - # all rows are active=1, so this is behaviorally identical - # to the full replace. - # archive_dropped: a rewind overwrites turns the user may - # not have meant to drop, and this write is the last step - # before they are gone — three reported incidents ended - # here with nothing to restore from (#70516, #80763, - # #82756). Soft-archiving keeps them on disk (active=0) and - # in the FTS index, so a mis-aimed cut is recoverable - # instead of terminal. The live transcript is unchanged. - db.replace_messages( - session["session_key"], - truncated, - active_only=True, - archive_dropped=True, - ) - except Exception as exc: - logger.error( - "prompt.submit: replace_messages failed for session %s " - "(ordinal=%d); refusing turn so memory and DB stay " - "aligned: %s", - sid, - ordinal, - exc, - exc_info=True, - ) - return _err( - rid, - 5008, - f"failed to persist history truncation: {exc}", - ) + try: + with _session_db(session) as db: + if db is not None: + # Keep main's durable row-id contract: this rewrite + # stamps fresh row ids onto the same surviving dicts, + # and the response below lets Desktop rebind them. + # ``truncated`` is carrier-aware, so a composite target + # becomes its pure hidden scaffold in the same atomic + # replace that archives the selected row and tail. + db.replace_messages( + session["session_key"], + truncated, + active_only=True, + archive_dropped=True, + ) + except Exception as exc: + logger.error( + "prompt.submit: replace_messages failed for session %s " + "(ordinal=%d); refusing turn so memory and DB stay " + "aligned: %s", + sid, + ordinal, + exc, + exc_info=True, + ) + return _err( + rid, + 5008, + f"failed to persist history truncation: {exc}", + ) session["history"] = truncated session["history_version"] = int(session.get("history_version", 0)) + 1 if db is not None: - # replace_messages re-inserted the surviving prefix as NEW rows - # and stamped fresh _row_id values onto these same dicts. - # Surface the surviving user-turn ids (in visible-user-ordinal - # order) so the client can rebind its cached rowId stamps — - # otherwise a second rewind targeting an older surviving turn - # sends the pre-rewind id and the fail-closed resolver refuses - # it with 4018 (#83202 review: consecutive-rewind staleness). - # Ordinal order matches the client's visible-user filter the - # same way truncate ordinals already do. Entries are None when - # a row somehow has no stamp — the client must drop its cached - # id for that turn rather than keep a stale one. survivor_user_row_ids = [ _message_row_id(truncated[i]) for i in _history_user_indices(truncated) diff --git a/tui_gateway/methods_session.py b/tui_gateway/methods_session.py index d0dc6d88d7..f139c2edef 100644 --- a/tui_gateway/methods_session.py +++ b/tui_gateway/methods_session.py @@ -2481,25 +2481,34 @@ def _(rid, params: dict) -> dict: ) removed = 0 with session["history_lock"]: - history = session.get("history", []) + if session.get("running"): + return _err( + rid, 4009, "session busy — /interrupt the current turn before /undo" + ) + history = _history_without_ephemeral_scaffolding( + session.get("history", []) + ) # Truncate from the last *real* user turn. Popping only trailing # assistant/tool then one user left timeline markers # (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. - from agent.context_compressor import is_user_originated_turn + from agent.context_compressor import user_originated_turn_view - last_user_idx = None - for i in range(len(history) - 1, -1, -1): - msg = history[i] - if is_user_originated_turn(msg): - last_user_idx = i - break - if last_user_idx is not None: - removed = len(history) - last_user_idx - del history[last_user_idx:] - session["history_version"] = int(session.get("history_version", 0)) + 1 + user_indices = [ + index + for index, message in enumerate(history) + if user_originated_turn_view(message) is not None + ] + if user_indices: + try: + _installed, _live_view, rewound_count = ( + _rewind_active_session_history(session, len(user_indices) - 1) + ) + removed = rewound_count + except Exception as exc: + return _err(rid, 5008, f"undo: {exc}") return _ok(rid, {"removed": removed}) diff --git a/tui_gateway/methods_tools.py b/tui_gateway/methods_tools.py index aae1f295ad..552970788c 100644 --- a/tui_gateway/methods_tools.py +++ b/tui_gateway/methods_tools.py @@ -706,41 +706,49 @@ def _(rid, params: dict) -> dict: return _err( rid, 4009, "session busy — /interrupt the current turn before /retry" ) - history = session.get("history", []) - if not history: - return _err(rid, 4018, "no previous user message to retry") - # Walk backwards to the last *real* user turn. Timeline bookkeeping - # rows (display_kind set) and compaction handoffs are durable - # role=user but must not count as user-originated asks — same - # predicate as CLI resume/count and the prompt.submit ordinal fix. - # Without this, /retry re-sends opaque markers (model_switch / - # async_delegation_complete / auto_continue / CONTEXT COMPACTION - # handoffs) and truncates only the marker instead of the failed - # exchange (#80622). - from agent.context_compressor import is_user_originated_turn + from agent.context_compressor import ( + history_before_user_originated_turn, + retryable_user_text, + user_originated_turn_view, + ) - last_user_idx = None - for i in range(len(history) - 1, -1, -1): - msg = history[i] - if is_user_originated_turn(msg): - last_user_idx = i - break - if last_user_idx is None: - return _err(rid, 4018, "no previous user message to retry") - content = history[last_user_idx].get("content", "") - if isinstance(content, list): - content = " ".join( - p.get("text", "") - for p in content - if isinstance(p, dict) and p.get("type") == "text" - ) - if not content: - return _err(rid, 4018, "last user message is empty") - # Truncate history: remove everything from the last user message onward - # (mirrors CLI retry_last() which strips the failed exchange) with session["history_lock"]: - session["history"] = history[:last_user_idx] - session["history_version"] = int(session.get("history_version", 0)) + 1 + if session.get("running"): + return _err( + rid, + 4009, + "session busy — /interrupt the current turn before /retry", + ) + if session.get("attached_images"): + return _err( + rid, + 4018, + "retry cannot safely reconstruct or combine attached media", + ) + history = _history_without_ephemeral_scaffolding( + session.get("history", []) + ) + user_indices = [ + index + for index, message in enumerate(history) + if user_originated_turn_view(message) is not None + ] + if not user_indices: + return _err(rid, 4018, "no previous user message to retry") + _prefix, live_view = history_before_user_originated_turn( + history, user_indices[-1] + ) + try: + content = retryable_user_text(live_view.get("content")) + except ValueError as exc: + return _err(rid, 4018, str(exc)) + try: + _active, durable_live_view, _rewound_count = ( + _rewind_active_session_history(session, len(user_indices) - 1) + ) + except Exception as exc: + return _err(rid, 5008, f"retry: failed to persist history: {exc}") + content = retryable_user_text(durable_live_view.get("content")) return _ok(rid, {"type": "send", "message": content}) if name == "steer": @@ -846,9 +854,6 @@ def _(rid, params: dict) -> dict: return _err( rid, 4009, "session busy — /interrupt the current turn before /undo" ) - db = _get_db() - if db is None: - return _db_unavailable_error(rid, code=5008) session_key = session.get("session_key", "") if not session_key: return _err(rid, 4001, "no session key for undo") @@ -862,35 +867,39 @@ def _(rid, params: dict) -> dict: return _err(rid, 4004, f"undo: invalid count {arg_str!r} — use /undo or /undo N") if n < 1: n = 1 - try: - recents = db.list_recent_user_messages(session_key, limit=max(n, 10)) - except Exception as e: - return _err(rid, 5008, f"undo: failed to load history: {e}") - if not recents: - return _err(rid, 4018, "no user messages to undo") - # recents[0] is the most-recent user turn; pick the Nth-from-last. - # If N exceeds the number of user turns, back up to the oldest. - target_idx = min(n - 1, len(recents) - 1) - target_id = recents[target_idx]["id"] - try: - result = db.rewind_to_message(session_key, target_id) - except ValueError as e: - return _err(rid, 4004, f"undo: {e}") - except Exception as e: - return _err(rid, 5008, f"undo: {e}") - # Reload the active-only transcript into the in-memory session - # history so subsequent turns see the truncated view. - # repair_alternation: this reload feeds LIVE REPLAY — session["history"] - # is the working conversation for subsequent turns, and a rewind that - # lands on a durable user;user pair would otherwise re-fire the - # pre-request repair on every request from here on. - try: - active = db.get_messages_as_conversation(session_key, repair_alternation=True) - except Exception: - active = [] + from agent.context_compressor import ( + user_originated_turn_view, + ) + from agent.message_content import flatten_message_text + with session["history_lock"]: - session["history"] = list(active) - session["history_version"] = int(session.get("history_version", 0)) + 1 + if session.get("running"): + return _err( + rid, + 4009, + "session busy — /interrupt the current turn before /undo", + ) + history = _history_without_ephemeral_scaffolding( + session.get("history", []) + ) + user_indices = [ + index + for index, message in enumerate(history) + if user_originated_turn_view(message) is not None + ] + if not user_indices: + return _err(rid, 4018, "no user messages to undo") + turns_undone = min(n, len(user_indices)) + target_position = len(user_indices) - turns_undone + try: + active, live_view, rewound_count = _rewind_active_session_history( + session, target_position + ) + except ValueError as exc: + return _err(rid, 4004, f"undo: {exc}") + except Exception as exc: + return _err(rid, 5008, f"undo: {exc}") + target_text = flatten_message_text(live_view.get("content")) # Notify memory providers — same hook /branch fires, plus the # rewound flag so providers caching per-turn document state # know to invalidate. See #6672 + #21910. @@ -917,18 +926,6 @@ def _(rid, params: dict) -> dict: agent._last_flushed_db_idx = len(active) except Exception: pass - target_msg = result.get("target_message") or {} - target_text = target_msg.get("content") or "" - if isinstance(target_text, list): - parts = [ - p.get("text", "") for p in target_text - if isinstance(p, dict) and p.get("type") == "text" - ] - target_text = "\n".join(t for t in parts if t) - if not isinstance(target_text, str): - target_text = "" - rewound_count = result.get("rewound_count", 0) - turns_undone = target_idx + 1 turn_word = "turn" if turns_undone == 1 else "turns" notice = ( f"↶ Undid {turns_undone} {turn_word} ({rewound_count} message(s)). " @@ -1297,27 +1294,28 @@ def _(rid, params: dict) -> dict: if result.get("success") and not file_path: removed = 0 with session["history_lock"]: - history = session.get("history", []) - # Truncate from the last *real* user turn. Same predicate - # as list_recent_user_messages / /undo / /retry — - # is_user_originated_turn also excludes compaction - # handoffs (durable role=user, sometimes without - # display_kind on legacy sessions; #80622). - from agent.context_compressor import is_user_originated_turn + history = _history_without_ephemeral_scaffolding( + session.get("history", []) + ) + from agent.context_compressor import user_originated_turn_view - last_user_idx = None - for i in range(len(history) - 1, -1, -1): - msg = history[i] - if is_user_originated_turn(msg): - last_user_idx = i - break - if last_user_idx is not None: - removed = len(history) - last_user_idx - del history[last_user_idx:] - if removed: - session["history_version"] = ( - int(session.get("history_version", 0)) + 1 - ) + user_indices = [ + index + for index, message in enumerate(history) + if user_originated_turn_view(message) is not None + ] + if user_indices: + try: + _active, _live_view, removed = ( + _rewind_active_session_history( + session, len(user_indices) - 1 + ) + ) + except Exception as exc: + raise RuntimeError( + "checkpoint restored, but session history rewind " + f"failed: {exc}" + ) from exc result["history_removed"] = removed return result diff --git a/tui_gateway/server.py b/tui_gateway/server.py index fc0c3f02f5..179ed6ca88 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -2980,6 +2980,140 @@ def _session_db(session: dict): db.close() +def _rewind_active_session_history( + session: dict, user_ordinal: int +) -> tuple[list[dict], dict, int]: + """Rewind one canonical user turn while retaining carrier scaffolding. + + The caller holds ``history_lock``. Persistent sessions archive the target + and tail; a composite carrier's own hidden handoff is inserted in that same + transaction. Memory is installed only after the durable commit and is + built from the already-validated prefix plus the returned scaffold row id, + so there is no fallible post-commit reload. + """ + 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, + ) + + 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 = [] + 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]") + 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 + + history = _history_without_ephemeral_scaffolding(session.get("history", [])) + user_indices = [ + index + for index, message in enumerate(history) + if user_originated_turn_view(message) 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") + durable = db.get_messages_as_conversation( + session_key, + include_row_ids=True, + ) + durable_user_indices = [ + index + for index, message in enumerate(durable) + if user_originated_turn_view(message) is not None + ] + 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" + ) + 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) + ] + 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: + replacement_id = result.get("replacement_message_id") + if not isinstance(replacement_id, int): + raise RuntimeError( + "rewind commit did not return the replacement scaffold id" + ) + durable_prefix[-1]["_row_id"] = replacement_id + durable_prefix[-1]["_db_persisted"] = True + installed[-1] = durable_prefix[-1] + live_view = durable_live_view + rewound_count = int(result.get("rewound_count", 0)) + persisted = True + + installed = [message.copy() for message in installed] + session["history"] = installed + session["history_version"] = int(session.get("history_version", 0)) + 1 + agent = session.get("agent") + 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 + if hasattr(agent, "_db_flush_scan_prefix"): + agent._db_flush_scan_prefix = installed[:] if persisted else None + return installed, live_view, rewound_count + + +def _history_without_ephemeral_scaffolding(history: list[dict]) -> list[dict]: + """Return the durable transcript shape without transient recovery rows.""" + from run_agent import _is_ephemeral_scaffolding + + return [ + message.copy() + for message in history + if not _is_ephemeral_scaffolding(message) + ] + + def _persist_session_git_meta(session: dict, cwd: str) -> None: """Resolve + persist a session's git branch / repo root WITHOUT blocking. @@ -7190,6 +7324,20 @@ 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