From 962189d9ea66326d0e40672cbb7fd6110452cc30 Mon Sep 17 00:00:00 2001 From: kshitijk4poor <82637225+kshitijk4poor@users.noreply.github.com> Date: Tue, 14 Jul 2026 00:35:56 +0530 Subject: [PATCH] fix(cli): clear stale persistence override before staging --- cli.py | 22 +++- run_agent.py | 8 +- tests/cli/test_cli_interrupt_ack_race.py | 131 +++++++++++++++++++++++ 3 files changed, 156 insertions(+), 5 deletions(-) diff --git a/cli.py b/cli.py index 45e32090febd..99b910b23fb8 100644 --- a/cli.py +++ b/cli.py @@ -12226,9 +12226,25 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin): # duplicate-prone intermediate snapshot to terminal-close persistence. if self.conversation_history is getattr(agent, "_session_messages", None): self.conversation_history = list(self.conversation_history) - staged_user_message = {"role": "user", "content": message} - agent._pending_cli_user_message = staged_user_message - self.conversation_history.append(staged_user_message) + # The prior turn's override applies only to its own user dict. Clear it + # before exposing the next staged input to close persistence; otherwise + # a shutdown before the worker prologue can write old API-local text as + # this new user message (#63766). + persist_lock = getattr(agent, "_session_persist_lock", None) + + def _stage_user_message() -> None: + agent._persist_user_message_idx = None + agent._persist_user_message_override = None + agent._persist_user_message_timestamp = None + staged_user_message = {"role": "user", "content": message} + agent._pending_cli_user_message = staged_user_message + self.conversation_history.append(staged_user_message) + + if persist_lock is None: + _stage_user_message() + else: + with persist_lock: + _stage_user_message() ChatConsole().print(f"[{_accent_hex()}]{'─' * 40}[/]") print(flush=True) diff --git a/run_agent.py b/run_agent.py index 89c4ce633733..cf8e63ff9b52 100644 --- a/run_agent.py +++ b/run_agent.py @@ -1764,7 +1764,11 @@ class AIAgent: from agent.agent_runtime_helpers import repair_message_sequence return repair_message_sequence(self, messages) - def _flush_messages_to_session_db(self, messages: List[Dict], conversation_history: List[Dict] = None): + def _flush_messages_to_session_db( + self, + messages: List[Dict], + conversation_history: Optional[List[Dict]] = None, + ): """Serialize direct and turn-boundary session flushes per agent.""" persist_lock = getattr(self, "_session_persist_lock", None) if persist_lock is None: @@ -1775,7 +1779,7 @@ class AIAgent: def _flush_messages_to_session_db_unlocked( self, messages: List[Dict], - conversation_history: List[Dict] = None, + conversation_history: Optional[List[Dict]] = None, ): """Persist any un-flushed messages to the SQLite session store. diff --git a/tests/cli/test_cli_interrupt_ack_race.py b/tests/cli/test_cli_interrupt_ack_race.py index dbcd7a15bd09..ac465cb20ca1 100644 --- a/tests/cli/test_cli_interrupt_ack_race.py +++ b/tests/cli/test_cli_interrupt_ack_race.py @@ -26,6 +26,7 @@ from __future__ import annotations import importlib import queue import sys +import threading import time from unittest.mock import MagicMock, patch @@ -231,3 +232,133 @@ def test_chat_persists_clean_input_when_a_queued_note_changes_api_message(): assert agent.captured is not None assert agent.captured["user_message"] == "[MODEL SWITCH NOTE]\n\nclean prompt" assert agent.captured["persist_user_message"] == "clean prompt" + + +def test_chat_clears_previous_turn_persistence_override_before_staging(): + """A close before the next worker starts cannot reuse a stale override.""" + cli = _make_cli() + + class _StagingAgent(_StubAgent): + def __init__(self, session_id): + super().__init__(session_id, turn_seconds=0) + self.staged_override = None + self.staged_message = None + self._session_messages = [] + self._persist_user_message_idx = 7 + self._persist_user_message_override = "previous clean prompt" + self._persist_user_message_timestamp = 123.0 + + def run_conversation(self, **kwargs): + self.staged_override = self._persist_user_message_override + self.staged_message = self._pending_cli_user_message + return { + "final_response": "done", + "messages": [{"role": "assistant", "content": "done"}], + "api_calls": 1, + "completed": True, + "partial": True, + "response_previewed": True, + } + + agent = _StagingAgent(cli.session_id) + cli.agent = agent + cli._interrupt_queue = queue.Queue() + cli._pending_input = queue.Queue() + + with patch.object(cli, "_ensure_runtime_credentials", return_value=True), \ + patch.object(cli, "_resolve_turn_agent_config", return_value={ + "signature": cli._active_agent_route_signature, + "model": None, "runtime": None, "request_overrides": None, + }), \ + patch.object(cli, "_init_agent", return_value=True): + cli.chat("new prompt") + + assert agent.staged_override is None + assert agent._persist_user_message_idx is None + assert agent._persist_user_message_timestamp is None + assert agent.staged_message == {"role": "user", "content": "new prompt"} + + +def test_chat_close_does_not_persist_previous_turn_override(tmp_path, monkeypatch): + """A close after input staging writes the new prompt, not old API-only text.""" + from hermes_state import SessionDB + from run_agent import AIAgent + + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + cli = _make_cli() + session_id = cli.session_id + db = SessionDB(db_path=tmp_path / "state.db") + db.create_session(session_id=session_id, source="cli") + prefix = [ + {"role": "user", "content": "old prompt"}, + {"role": "assistant", "content": "old answer"}, + ] + for message in prefix: + db.append_message( + session_id=session_id, + role=message["role"], + content=message["content"], + ) + + agent = object.__new__(AIAgent) + agent._session_db = db + agent._session_db_created = True + agent.session_id = session_id + agent.platform = "cli" + agent.model = "test-model" + agent._session_messages = [] + agent._last_flushed_db_idx = 0 + agent._flushed_db_message_ids = set() + agent._flushed_db_message_session_id = None + agent._persist_disabled = False + agent._cached_system_prompt = "test system prompt" + agent._session_init_model_config = None + agent._parent_session_id = None + agent._session_json_enabled = False + agent._pending_cli_user_message = None + agent._session_persist_lock = threading.RLock() + agent._persist_user_message_idx = len(prefix) + agent._persist_user_message_override = "previous clean prompt" + agent._persist_user_message_timestamp = 123.0 + agent._active_children = [] + agent._interrupt_requested = False + entered = threading.Event() + release = threading.Event() + + def _block_run(**_kwargs): + entered.set() + assert release.wait(timeout=5) + return { + "final_response": "done", + "messages": prefix + [{"role": "assistant", "content": "done"}], + "api_calls": 1, + "completed": True, + "partial": True, + "response_previewed": True, + } + + agent.run_conversation = _block_run + cli.agent = agent + cli.conversation_history = list(prefix) + cli._interrupt_queue = queue.Queue() + cli._pending_input = queue.Queue() + + with patch.object(cli, "_ensure_runtime_credentials", return_value=True), \ + patch.object(cli, "_resolve_turn_agent_config", return_value={ + "signature": cli._active_agent_route_signature, + "model": None, "runtime": None, "request_overrides": None, + }), \ + patch.object(cli, "_init_agent", return_value=True): + chat_thread = threading.Thread(target=lambda: cli.chat("new prompt")) + chat_thread.start() + assert entered.wait(timeout=5) + cli._persist_active_session_before_close() + release.set() + chat_thread.join(timeout=10) + + assert not chat_thread.is_alive() + assert [m["content"] for m in db.get_messages_as_conversation(session_id)] == [ + "old prompt", + "old answer", + "new prompt", + ]