From 32bdc67e104934ca7324df430437e0cd839e01c3 Mon Sep 17 00:00:00 2001 From: kshitijk4poor <82637225+kshitijk4poor@users.noreply.github.com> Date: Tue, 14 Jul 2026 01:24:46 +0530 Subject: [PATCH] fix(cli): snapshot close state under staging lock --- cli.py | 112 +++++++++++---------- tests/cli/test_cli_interrupt_ack_race.py | 122 +++++++++++++++++++++++ 2 files changed, 180 insertions(+), 54 deletions(-) diff --git a/cli.py b/cli.py index 99b910b23fb8..92010a1000e8 100644 --- a/cli.py +++ b/cli.py @@ -12857,60 +12857,64 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin): if not agent or not hasattr(agent, "_persist_session"): return - messages = getattr(agent, "_session_messages", None) - pending_cli_message = getattr(agent, "_pending_cli_user_message", None) - if not isinstance(messages, list): - messages = getattr(self, "conversation_history", None) - if not isinstance(messages, list): - return - if isinstance(pending_cli_message, dict) and not any( - message is pending_cli_message for message in messages - ): - # The UI has accepted a new input but the worker still exposes its - # prior snapshot. Include only that staged dict; the baseline below - # keeps any durable resumed prefix from being re-appended. - messages = [*messages, pending_cli_message] - if not messages: - return - - # A normal turn builds a new list that reuses the resumed-history dicts. - # Keep that CLI history as the baseline so a signal between assigning - # ``_session_messages`` and the turn's DB flush cannot append its durable - # prefix a second time. Once the CLI takes the turn result, however, both - # names can point at the same live list; passing that alias would mark an - # unflushed tail durable without writing it. Marker-only persistence is - # correct only in that alias case. - conversation_history = getattr(self, "conversation_history", None) - pending_cli_message = getattr(agent, "_pending_cli_user_message", None) - if ( - isinstance(conversation_history, list) - and conversation_history - and conversation_history[-1] is pending_cli_message - ): - # The UI accepted this user message before the agent finished its - # early persistence. Its dict can already be in ``messages`` but is - # not durable yet, so exclude it from the resumed-history baseline. - conversation_history = conversation_history[:-1] - elif not isinstance(conversation_history, list) or conversation_history is messages: - conversation_history = None - - # A first-turn close can arrive before the worker builds its cached - # prompt. Build or restore it before the DB row is created so the - # durable transcript never leaves a NULL system_prompt cache entry. - if getattr(agent, "_cached_system_prompt", None) is None: - try: - from agent.conversation_loop import _restore_or_build_system_prompt - - _restore_or_build_system_prompt(agent, None, conversation_history) - except Exception: - logger.debug("Could not build system prompt during CLI close", exc_info=True) - return - if getattr(agent, "_cached_system_prompt", None) is None: - return - persist_lock = getattr(agent, "_session_persist_lock", None) - def _ensure_and_persist() -> None: + def _snapshot_and_persist() -> None: + # This snapshot must share the staging lock with ``chat()``. Without + # it, close can retain a mutable history baseline just before chat + # appends its pending dict; the later flush then mistakes that dict + # for durable history and stamps it without writing a row (#63766). + messages = getattr(agent, "_session_messages", None) + pending_cli_message = getattr(agent, "_pending_cli_user_message", None) + if not isinstance(messages, list): + messages = getattr(self, "conversation_history", None) + if not isinstance(messages, list): + return + if isinstance(pending_cli_message, dict) and not any( + message is pending_cli_message for message in messages + ): + # The UI has accepted a new input but the worker still exposes its + # prior snapshot. Include only that staged dict; the baseline below + # keeps any durable resumed prefix from being re-appended. + messages = [*messages, pending_cli_message] + if not messages: + return + + # A normal turn builds a new list that reuses the resumed-history dicts. + # Keep that CLI history as the baseline so a signal between assigning + # ``_session_messages`` and the turn's DB flush cannot append its durable + # prefix a second time. Once the CLI takes the turn result, however, both + # names can point at the same live list; passing that alias would mark an + # unflushed tail durable without writing it. Marker-only persistence is + # correct only in that alias case. + conversation_history = getattr(self, "conversation_history", None) + pending_cli_message = getattr(agent, "_pending_cli_user_message", None) + if ( + isinstance(conversation_history, list) + and conversation_history + and conversation_history[-1] is pending_cli_message + ): + # The UI accepted this user message before the agent finished its + # early persistence. Its dict can already be in ``messages`` but is + # not durable yet, so exclude it from the resumed-history baseline. + conversation_history = conversation_history[:-1] + elif not isinstance(conversation_history, list) or conversation_history is messages: + conversation_history = None + + # A first-turn close can arrive before the worker builds its cached + # prompt. Build or restore it before the DB row is created so the + # durable transcript never leaves a NULL system_prompt cache entry. + if getattr(agent, "_cached_system_prompt", None) is None: + try: + from agent.conversation_loop import _restore_or_build_system_prompt + + _restore_or_build_system_prompt(agent, None, conversation_history) + except Exception: + logger.debug("Could not build system prompt during CLI close", exc_info=True) + return + if getattr(agent, "_cached_system_prompt", None) is None: + return + agent._ensure_db_session() agent._persist_session(messages, conversation_history) if getattr(agent, "session_id", None): @@ -12918,10 +12922,10 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin): try: if persist_lock is None: - _ensure_and_persist() + _snapshot_and_persist() else: with persist_lock: - _ensure_and_persist() + _snapshot_and_persist() except (Exception, KeyboardInterrupt) as e: logger.debug("Could not persist active CLI session before close: %s", e) diff --git a/tests/cli/test_cli_interrupt_ack_race.py b/tests/cli/test_cli_interrupt_ack_race.py index ac465cb20ca1..c4e8036d67dd 100644 --- a/tests/cli/test_cli_interrupt_ack_race.py +++ b/tests/cli/test_cli_interrupt_ack_race.py @@ -362,3 +362,125 @@ def test_chat_close_does_not_persist_previous_turn_override(tmp_path, monkeypatc "old answer", "new prompt", ] + + +def test_close_waits_for_atomic_cli_staging_before_snapshot(tmp_path, monkeypatch): + """Close cannot retain the mutable pre-append history as its DB baseline.""" + 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" + # Deliberately distinct from CLI history: this is the normal pre-worker + # state that used to let close retain the wrong mutable baseline. + agent._session_messages = list(prefix) + 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 = None + agent._persist_user_message_override = None + agent._persist_user_message_timestamp = None + agent._active_children = [] + agent._interrupt_requested = False + + staging_entered = threading.Event() + release_staging = threading.Event() + run_entered = threading.Event() + release_run = threading.Event() + + class _BlockingHistory(list): + def __init__(self, values): + super().__init__(values) + self._block_next_append = True + + def append(self, value): + if self._block_next_append: + self._block_next_append = False + staging_entered.set() + assert release_staging.wait(timeout=5) + return super().append(value) + + def _block_run(**_kwargs): + run_entered.set() + assert release_run.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 = _BlockingHistory(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 staging_entered.wait(timeout=5) + + close_started = threading.Event() + close_finished = threading.Event() + + def _close(): + close_started.set() + cli._persist_active_session_before_close() + close_finished.set() + + close_thread = threading.Thread(target=_close) + close_thread.start() + assert close_started.wait(timeout=5) + # The close snapshot must wait for the locked pending-pointer/history + # handoff; otherwise the subsequent append poisons its DB baseline. + assert not close_finished.wait(timeout=0.1) + + release_staging.set() + assert run_entered.wait(timeout=5) + assert close_finished.wait(timeout=5) + release_run.set() + chat_thread.join(timeout=10) + close_thread.join(timeout=10) + + assert not chat_thread.is_alive() + assert not close_thread.is_alive() + assert [m["content"] for m in db.get_messages_as_conversation(session_id)] == [ + "old prompt", + "old answer", + "new prompt", + ]