fix(agent): preserve live turns in compaction carriers

This commit is contained in:
poisdahl
2026-08-13 21:38:30 +02:00
parent c692312704
commit 3e5e4c5d20
21 changed files with 2688 additions and 333 deletions
+217 -49
View File
@@ -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: