diff --git a/tests/cli/test_cli_interrupt_ack_race.py b/tests/cli/test_cli_interrupt_ack_race.py index c4e8036d67dd..0e2c21b6059c 100644 --- a/tests/cli/test_cli_interrupt_ack_race.py +++ b/tests/cli/test_cli_interrupt_ack_race.py @@ -28,6 +28,7 @@ import queue import sys import threading import time +import types from unittest.mock import MagicMock, patch @@ -234,6 +235,180 @@ def test_chat_persists_clean_input_when_a_queued_note_changes_api_message(): assert agent.captured["persist_user_message"] == "clean prompt" +def test_chat_preserves_clean_multimodal_input_when_note_changes_api_message(): + """A queued note forwards original native parts as the persistence override.""" + cli = _make_cli() + + class _NoteAgent(_StubAgent): + def __init__(self, session_id): + super().__init__(session_id, turn_seconds=0) + self.captured = None + + def run_conversation(self, **kwargs): + self.captured = kwargs + return { + "final_response": "done", + "messages": [{"role": "assistant", "content": "done"}], + "api_calls": 1, + "completed": True, + "partial": True, + "response_previewed": True, + } + + clean_parts = [ + {"type": "text", "text": "Describe this screenshot"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}, + ] + agent = _NoteAgent(cli.session_id) + cli.agent = agent + cli._interrupt_queue = queue.Queue() + cli._pending_input = queue.Queue() + cli._pending_model_switch_note = "[MODEL SWITCH NOTE]" + + 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(clean_parts) + + assert agent.captured is not None + assert agent.captured["persist_user_message"] == clean_parts + assert agent.captured["persist_user_message"] is not agent.captured["user_message"] + api_parts = agent.captured["user_message"] + assert api_parts[0]["text"] == "[MODEL SWITCH NOTE]\n\nDescribe this screenshot" + assert api_parts[1] == clean_parts[1] + + +def test_chat_multimodal_note_persists_clean_input_once(tmp_path, monkeypatch): + """The real CLI-to-agent path stores clean image parts, never the queued note.""" + 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") + + 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.provider = "test" + agent.base_url = "" + agent.api_key = "" + agent.api_mode = "chat_completions" + 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 = None + agent._persist_user_message_override = None + agent._persist_user_message_timestamp = None + agent._active_children = [] + agent._interrupt_requested = False + + clean_parts = [ + {"type": "text", "text": "Describe this screenshot"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,AAAA"}}, + ] + captured = {} + + def _realish_run(**kwargs): + captured.update(kwargs) + # Drive production turn setup and the real SQLite persistence seam, + # then return a normal CLI result without starting a provider loop. + from agent.turn_context import build_turn_context + + agent.quiet_mode = True + agent.max_iterations = 1 + agent.tools = [] + agent.valid_tool_names = set() + agent.enabled_toolsets = None + agent.disabled_toolsets = None + agent._skip_mcp_refresh = True + agent.compression_enabled = False + agent.context_compressor = types.SimpleNamespace(protect_first_n=2, protect_last_n=2) + agent._memory_store = None + agent._memory_manager = None + agent._memory_nudge_interval = 0 + agent._turns_since_memory = 0 + agent._user_turn_count = 0 + agent._todo_store = types.SimpleNamespace(has_items=lambda: True) + agent._tool_guardrails = types.SimpleNamespace(reset_for_turn=lambda: None) + agent._compression_warning = None + agent._memory_write_origin = "assistant_tool" + agent._stream_context_scrubber = None + agent._stream_think_scrubber = None + agent._restore_primary_runtime = lambda: None + agent._cleanup_dead_connections = lambda: False + agent._emit_status = lambda _message: None + agent._replay_compression_warning = lambda: None + agent._hydrate_todo_store = lambda *_args: None + agent._safe_print = lambda *_args: None + + context = build_turn_context( + agent, + kwargs["user_message"], + None, + kwargs["conversation_history"], + kwargs["task_id"], + None, + kwargs["persist_user_message"], + None, + restore_or_build_system_prompt=lambda *_args: None, + install_safe_stdio=lambda: None, + sanitize_surrogates=lambda value: value, + summarize_user_message_for_log=lambda value: ( + value if isinstance(value, str) else "[multimodal test message]" + ), + set_session_context=lambda _session_id: None, + set_current_write_origin=lambda _origin: None, + ra=lambda: types.SimpleNamespace(_set_interrupt=lambda *_args: None), + ) + agent._apply_persist_user_message_override(context.messages) + agent._persist_session(context.messages, kwargs["conversation_history"]) + return { + "final_response": "done", + "messages": context.messages + [{"role": "assistant", "content": "done"}], + "api_calls": 1, + "completed": True, + "partial": True, + "response_previewed": True, + } + + agent.run_conversation = _realish_run + cli.agent = agent + cli._interrupt_queue = queue.Queue() + cli._pending_input = queue.Queue() + cli._pending_model_switch_note = "[MODEL SWITCH NOTE]" + + 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(clean_parts) + + assert captured["persist_user_message"] == clean_parts + assert captured["user_message"][0]["text"] == "[MODEL SWITCH NOTE]\n\nDescribe this screenshot" + assert [m["content"] for m in db.get_messages_as_conversation(session_id)] == [ + "Describe this screenshot\n[screenshot]" + ] + + 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()