diff --git a/cecli/args.py b/cecli/args.py index 317523944fe..6d6b4ff3619 100644 --- a/cecli/args.py +++ b/cecli/args.py @@ -598,6 +598,12 @@ def get_parser(default_config_files, git_root): default=True, help="Enable/disable streaming responses (default: True)", ) + group.add_argument( + "--spinner", + action=argparse.BooleanOptionalAction, + default=True, + help="Enable/disable the spinner while waiting for LLM responses (default: True)", + ) group.add_argument( "--user-input-color", default="#00cc00", diff --git a/cecli/coders/agent_coder.py b/cecli/coders/agent_coder.py index 943f6b1c538..62c590f9adb 100644 --- a/cecli/coders/agent_coder.py +++ b/cecli/coders/agent_coder.py @@ -305,6 +305,25 @@ async def _exec_async(): call_result = await litellm.experimental_mcp_client.call_openai_tool( session=session, openai_tool=tool_call_dict ) + except Exception as e: + if server.is_session_expired_error(e): + try: + session = await server.reconnect() + call_result = await litellm.experimental_mcp_client.call_openai_tool( + session=session, openai_tool=tool_call_dict + ) + except Exception as retry_exc: + self.io.tool_warning( + f"Executing {tool_name} on {server.name} failed after reconnect:\n" + f"Error: {retry_exc}" + ) + return f"Error executing tool call {tool_name}: {retry_exc}" + else: + self.io.tool_warning( + f"Executing {tool_name} on {server.name} failed:\nError: {e}" + ) + return f"Error executing tool call {tool_name}: {e}" + try: content_parts = [] if call_result.content: for item in call_result.content: diff --git a/cecli/coders/base_coder.py b/cecli/coders/base_coder.py index 2633bc319e1..629ec23d987 100755 --- a/cecli/coders/base_coder.py +++ b/cecli/coders/base_coder.py @@ -221,6 +221,7 @@ def total_cached_tokens(self, value): message_tokens_sent = 0 message_tokens_received = 0 message_cached_tokens = 0 + message_cost_deferred = None add_cache_headers = False cache_warming_thread = None num_cache_warming_pings = 0 @@ -1497,6 +1498,10 @@ async def _run_linear(self, with_message=None, preproc=True): self.show_announcements() self.suppress_announcements_for_next_prompt = True + if self.message_cost_deferred and not self.io.spinner_active: + self.io.tool_output(self.message_cost_deferred) + self.message_cost_deferred = None + await self.io.recreate_input() await self.io.input_task user_message = self.io.input_task.result() @@ -1645,6 +1650,10 @@ async def input_task(self, preproc): self.show_announcements() self.suppress_announcements_for_next_prompt = True + if self.message_cost_deferred and not self.io.spinner_active: + self.io.tool_output(self.message_cost_deferred) + self.message_cost_deferred = None + # Stop spinner before showing announcements or getting input self.io.stop_spinner() self.copy_context() @@ -2512,7 +2521,11 @@ async def format_in_executor(): if not self.tui: spinner_text += f" • ${self.format_cost(self.total_cost)} session" - self.io.start_spinner(spinner_text, coder_uuid=getattr(self, "uuid", None)) + if self.io.spinner_active: + self.io.start_spinner(spinner_text, coder_uuid=getattr(self, "uuid", None)) + else: + self.message_cost_deferred = spinner_text + if self.stream: self.mdstream = True else: @@ -2635,6 +2648,10 @@ async def format_in_executor(): # Ensure any waiting spinner is stopped self.io.start_spinner("Processing Answer...", coder_uuid=getattr(self, "uuid", None)) + + if not self.io.spinner_active: + self.partial_response_content = self.get_multi_response_content_in_progress(True) + self.remove_reasoning_content() self.multi_response_content = "" @@ -2980,12 +2997,22 @@ async def _execute_mcp_tools(self, server, tool_calls): continue async def do_tool_call(): + nonlocal session from litellm import experimental_mcp_client - return await experimental_mcp_client.call_openai_tool( - session=session, - openai_tool=new_tool_call, - ) + try: + return await experimental_mcp_client.call_openai_tool( + session=session, + openai_tool=new_tool_call, + ) + except Exception as e: + if server.is_session_expired_error(e): + session = await server.reconnect() + return await experimental_mcp_client.call_openai_tool( + session=session, + openai_tool=new_tool_call, + ) + raise call_result, interrupted = await coroutines.interruptible( do_tool_call(), self.interrupt_event diff --git a/cecli/interruptible_input.py b/cecli/interruptible_input.py index e93eb96fa24..52abfb86c84 100644 --- a/cecli/interruptible_input.py +++ b/cecli/interruptible_input.py @@ -17,7 +17,14 @@ def __init__(self): raise RuntimeError("InterruptibleInput is Unix-only (requires selectable stdin).") self._cancel = threading.Event() - self._sel = selectors.DefaultSelector() + + # The default selector (Kqueue on macOS, Epoll on Linux) cannot + # handle pipe-based stdin (e.g. when running inside Emacs comint-mode). + # Fall back to SelectSelector which works with any fd that supports select(). + if not sys.stdin.isatty(): + self._sel = selectors.SelectSelector() + else: + self._sel = selectors.DefaultSelector() # self-pipe to wake up select() from interrupt() self._r, self._w = os.pipe() diff --git a/cecli/io.py b/cecli/io.py index 04438faaad9..9d03c9be88a 100644 --- a/cecli/io.py +++ b/cecli/io.py @@ -369,6 +369,7 @@ def __init__( notifications_command=None, notification_bell=False, verbose=False, + show_spinner=True, ): self.console = Console() self.pretty = pretty @@ -499,6 +500,7 @@ def __init__( fancy_input = False # Spinner state + self.spinner_active = show_spinner self.spinner_running = False self.spinner_text = "" self.last_spinner_text = "" @@ -507,7 +509,7 @@ def __init__( self.spinner_last_frame_index = 0 self.unicode_palette = "░█" self.fallback_spinner = None - self.fallback_spinner_enabled = True + self.fallback_spinner_enabled = show_spinner self.interruptible_input = None @@ -569,7 +571,12 @@ def start_spinner(self, text, update_last_text=True, **kwargs): """Start the spinner.""" self.stop_spinner() + if not self.spinner_active: + return + if self.prompt_session: + if not self.fallback_spinner_enabled: + return self.spinner_running = True self.spinner_text = text self.spinner_frame_index = self.spinner_last_frame_index @@ -582,9 +589,15 @@ def start_spinner(self, text, update_last_text=True, **kwargs): self.fallback_spinner.step() def update_spinner(self, text): + if not self.spinner_active: + return + self.spinner_text = text def update_spinner_suffix(self, text=None): + if not self.spinner_active: + return + if text: self.spinner_suffix = f" • {text[:16].strip()}" else: diff --git a/cecli/main.py b/cecli/main.py index aea0c6b8690..aeaf49f1acd 100644 --- a/cecli/main.py +++ b/cecli/main.py @@ -40,6 +40,20 @@ if sys.platform == "win32": if hasattr(asyncio, "set_event_loop_policy"): asyncio.set_event_loop_policy(asyncio.WindowsSelectorEventLoopPolicy()) +elif sys.platform == "darwin": + # The default KqueueSelector cannot handle pipe-based stdin + # (e.g. when running inside Emacs comint-mode). Fall back to + # SelectSelector which works with any file descriptor that supports select(). + import selectors + + if not sys.stdin.isatty(): + _original_event_loop_policy = asyncio.DefaultEventLoopPolicy + + class _SelectSelectorPolicy(asyncio.DefaultEventLoopPolicy): + def new_event_loop(self): + return asyncio.SelectorEventLoop(selectors.SelectSelector()) + + asyncio.set_event_loop_policy(_SelectSelectorPolicy()) from prompt_toolkit.enums import EditingMode from .dump import dump # noqa @@ -708,6 +722,7 @@ def get_io(pretty): notifications_command=args.notifications_command, notification_bell=args.notification_bell, verbose=args.verbose, + show_spinner=args.spinner, ) validate_tui_args(args) diff --git a/cecli/mcp/manager.py b/cecli/mcp/manager.py index f5211cd0df8..bcb4d56e7fe 100644 --- a/cecli/mcp/manager.py +++ b/cecli/mcp/manager.py @@ -167,7 +167,7 @@ async def connect_server(self, name: str) -> bool: # When io is None (e.g., during from_servers before IO is assigned), # _log_warning and _log_error silently return — retries still happen # but with no user-visible feedback. This is intentional. - max_retries = 3 + max_retries = 3 if server.name != "unnamed-server" else 1 delay = 1.0 backoff = 2.0 max_delay = 30.0 @@ -185,11 +185,12 @@ async def connect_server(self, name: str) -> bool: except asyncio.CancelledError: raise except Exception as e: - if attempt < max_retries: + if attempt < max_retries and server.name != "unnamed-server": self._log_warning( f"Connection attempt {attempt} failed for {name}, " f"retrying in {delay}s... ({e})" ) + await asyncio.sleep(delay) delay = min(delay * backoff, max_delay) else: diff --git a/cecli/mcp/server.py b/cecli/mcp/server.py index bb92c1473dd..4d5240256e4 100644 --- a/cecli/mcp/server.py +++ b/cecli/mcp/server.py @@ -130,6 +130,48 @@ async def disconnect(self): finally: self.session = None + async def reconnect(self): + """Disconnect and reconnect, establishing a fresh session. + + Used when the server has invalidated the current session (e.g., after + a server restart), as indicated by an HTTP 404 response per the MCP + protocol specification. + + Returns: + ClientSession: The new active session + """ + if self.io: + self.io.tool_warning(f"MCP session expired for {self.name}, reconnecting...") + await self.disconnect() + self.exit_stack = AsyncExitStack() + return await self.connect() + + @staticmethod + def is_session_expired_error(exc): + """Check if an exception indicates an expired MCP session (HTTP 404). + + Per the MCP specification, when a server terminates a session it + responds with HTTP 404 Not Found. The client MUST then start a new + session by sending a new InitializeRequest. + + Args: + exc: The exception to check + + Returns: + bool: True if the error indicates a 404 session expiry + """ + import httpx + + if isinstance(exc, httpx.HTTPStatusError) and exc.response.status_code == 404: + return True + + # Some transports wrap the status in the exception message + exc_str = str(exc).lower() + if "404" in exc_str and ("session" in exc_str or "not found" in exc_str): + return True + + return False + class HttpBasedMcpServer(McpServer): """Base class for HTTP-based MCP servers (HTTP streaming and SSE).""" diff --git a/cecli/repo.py b/cecli/repo.py index 4b42387257e..baec1592a99 100644 --- a/cecli/repo.py +++ b/cecli/repo.py @@ -486,6 +486,7 @@ async def get_commit_message(self, diffs, context, user_language=None): commit_message = None for model in self.models: spinner_text = f"Generating commit message with {model.name}\n" + self.io.start_spinner(spinner_text, update_last_text=False) if model.system_prompt_prefix: diff --git a/cecli/website/docs/config/conf.md b/cecli/website/docs/config/conf.md index efcb199eb1e..ad3c1a8578c 100644 --- a/cecli/website/docs/config/conf.md +++ b/cecli/website/docs/config/conf.md @@ -236,6 +236,9 @@ cog.outl("```") ## Enable/disable streaming responses (default: True) #stream: true +## Enable/disable the spinner while waiting for LLM responses (default: True) +#spinner: true + ## Set the color for user input (default: #00cc00) #user-input-color: "#00cc00" diff --git a/cecli/website/docs/config/options.md b/cecli/website/docs/config/options.md index f462120ede7..d196c44adcb 100644 --- a/cecli/website/docs/config/options.md +++ b/cecli/website/docs/config/options.md @@ -351,6 +351,14 @@ Aliases: - `--stream` - `--no-stream` +### `--spinner` +Enable/disable the spinner while waiting for LLM responses (default: True) +Default: True +Environment variable: `CECLI_SPINNER` +Aliases: + - `--spinner` + - `--no-spinner` + ### `--user-input-color VALUE` Set the color for user input (default: #00cc00) Default: #00cc00 diff --git a/tests/basic/test_select_selector.py b/tests/basic/test_select_selector.py new file mode 100644 index 00000000000..9a81d3b1fda --- /dev/null +++ b/tests/basic/test_select_selector.py @@ -0,0 +1,140 @@ +"""Tests for SelectSelector fallback used when stdin is not a TTY. + +Covers: +- cecli/interruptible_input.py: SelectSelector vs DefaultSelector choice +- cecli/main.py: _SelectSelectorPolicy on macOS when stdin is not a TTY +""" + +import asyncio +import os +import selectors +import sys +from unittest import mock + +import pytest + +from cecli.interruptible_input import InterruptibleInput + +# --------------------------------------------------------------------------- +# InterruptibleInput selector tests +# --------------------------------------------------------------------------- + + +class TestInterruptibleInputSelector: + """InterruptibleInput should pick SelectSelector for non-TTY stdin.""" + + @pytest.mark.skipif(os.name == "nt", reason="Unix-only") + def test_uses_select_selector_when_not_a_tty(self): + with mock.patch.object(sys.stdin, "isatty", return_value=False): + obj = InterruptibleInput() + try: + assert isinstance(obj._sel, selectors.SelectSelector) + finally: + obj.close() + + @pytest.mark.skipif(os.name == "nt", reason="Unix-only") + def test_uses_default_selector_when_tty(self): + with mock.patch.object(sys.stdin, "isatty", return_value=True): + obj = InterruptibleInput() + try: + assert isinstance(obj._sel, selectors.DefaultSelector) + finally: + obj.close() + + @pytest.mark.skipif(os.name == "nt", reason="Unix-only") + def test_selector_registers_wakeup_pipe(self): + with mock.patch.object(sys.stdin, "isatty", return_value=False): + obj = InterruptibleInput() + try: + # The wakeup read-end fd should be registered + key = obj._sel.get_key(obj._r) + assert key.data == "__wakeup__" + assert key.events & selectors.EVENT_READ + finally: + obj.close() + + @pytest.mark.skipif(os.name == "nt", reason="Unix-only") + def test_close_is_safe_to_call_twice(self): + with mock.patch.object(sys.stdin, "isatty", return_value=False): + obj = InterruptibleInput() + obj.close() + # Second close should not raise + obj.close() + + @pytest.mark.skipif(os.name == "nt", reason="Unix-only") + def test_interrupt_sets_cancel_and_wakes_selector(self): + with mock.patch.object(sys.stdin, "isatty", return_value=False): + obj = InterruptibleInput() + try: + obj.interrupt() + assert obj._cancel.is_set() + # The wakeup pipe should have data + data = os.read(obj._r, 1024) + assert len(data) > 0 + finally: + obj.close() + + @pytest.mark.skipif(os.name == "nt", reason="Unix-only") + def test_input_raises_interrupted_when_cancelled_before_call(self): + with mock.patch.object(sys.stdin, "isatty", return_value=False): + obj = InterruptibleInput() + try: + obj.interrupt() + with pytest.raises(InterruptedError, match="Input interrupted"): + obj.input("") + finally: + obj.close() + + @pytest.mark.skipif(os.name == "nt", reason="Unix-only") + def test_raises_on_windows(self): + with mock.patch("os.name", "nt"): + with pytest.raises(RuntimeError, match="Unix-only"): + InterruptibleInput() + + +# --------------------------------------------------------------------------- +# macOS _SelectSelectorPolicy tests +# --------------------------------------------------------------------------- + + +class TestSelectSelectorPolicyMacOS: + """On macOS with non-TTY stdin, the event loop should use SelectSelector.""" + + @pytest.mark.skipif(sys.platform != "darwin", reason="macOS-only policy") + def test_policy_uses_select_selector_on_macos_non_tty(self): + """When stdin is not a TTY on macOS, the patched policy should + produce a SelectorEventLoop backed by SelectSelector.""" + from cecli.main import _SelectSelectorPolicy + + policy = _SelectSelectorPolicy() + loop = policy.new_event_loop() + try: + selector = loop._selector + assert isinstance(selector, selectors.SelectSelector) + finally: + loop.close() + + def test_main_module_sets_policy_on_darwin_non_tty(self): + """Simulate importing the selector-policy block on macOS with piped stdin.""" + select_selector_cls = selectors.SelectSelector + + # Build a mini _SelectSelectorPolicy the same way main.py does + class _SelectSelectorPolicy(asyncio.DefaultEventLoopPolicy): + def new_event_loop(self): + return asyncio.SelectorEventLoop(select_selector_cls()) + + policy = _SelectSelectorPolicy() + loop = policy.new_event_loop() + try: + assert isinstance(loop._selector, selectors.SelectSelector) + finally: + loop.close() + + def test_default_policy_not_changed_when_tty(self): + """When stdin IS a TTY, the default event loop policy should remain.""" + original_policy = asyncio.get_event_loop_policy() + with mock.patch.object(sys.stdin, "isatty", return_value=True): + # The policy should be whatever the system default is, + # not our custom _SelectSelectorPolicy + current = asyncio.get_event_loop_policy() + assert current is original_policy diff --git a/tests/basic/test_spinner.py b/tests/basic/test_spinner.py new file mode 100644 index 00000000000..e66965bafd8 --- /dev/null +++ b/tests/basic/test_spinner.py @@ -0,0 +1,86 @@ +"""Tests for the --spinner / --no-spinner CLI option.""" + +from unittest.mock import MagicMock + +import pytest + + +@pytest.fixture +def mock_io(): + io = MagicMock() + io.last_spinner_text = "" + return io + + +@pytest.fixture +def mock_model(): + model = MagicMock() + model.name = "test-model" + model.system_prompt_prefix = None + model.send_completion = MagicMock( + return_value=MagicMock(choices=[MagicMock(message=MagicMock(content="test commit"))]) + ) + model.token_count = MagicMock(return_value=10) + model.info = {"max_input_tokens": 100000} + model.simple_send_with_retries = MagicMock(return_value="test commit") + + async def _async_simple_send(*args, **kwargs): + return "test commit" + + model.simple_send_with_retries = _async_simple_send + return model + + +class TestSpinnerArgParsing: + """Tests that argparse correctly handles --spinner / --no-spinner.""" + + def test_spinner_default_is_true(self): + """The default value for --spinner should be True.""" + from cecli.args import get_parser + + parser = get_parser(default_config_files=[], git_root=None) + args = parser.parse_args([]) + assert args.spinner is True + + def test_spinner_flag_sets_true(self): + """Passing --spinner explicitly sets spinner to True.""" + from cecli.args import get_parser + + parser = get_parser(default_config_files=[], git_root=None) + args = parser.parse_args(["--spinner"]) + assert args.spinner is True + + def test_no_spinner_flag_sets_false(self): + """Passing --no-spinner sets spinner to False.""" + from cecli.args import get_parser + + parser = get_parser(default_config_files=[], git_root=None) + args = parser.parse_args(["--no-spinner"]) + assert args.spinner is False + + +class TestIOSpinnerGating: + """Tests that InputOutput.start_spinner respects show_spinner=False.""" + + def test_io_show_spinner_false_disables_fallback_spinner(self): + """When show_spinner=False, fallback_spinner_enabled is False.""" + from cecli.io import InputOutput + + io = InputOutput(pretty=False, show_spinner=False) + assert io.fallback_spinner_enabled is False + + def test_io_show_spinner_true_by_default(self): + """By default, fallback_spinner_enabled is True.""" + from cecli.io import InputOutput + + io = InputOutput(pretty=False) + assert io.fallback_spinner_enabled is True + + def test_io_start_spinner_noop_when_disabled(self): + """start_spinner should not create a fallback spinner when show_spinner=False.""" + from cecli.io import InputOutput + + io = InputOutput(pretty=False, show_spinner=False) + io.start_spinner("Awaiting Confirmation...") + assert io.fallback_spinner is None + assert io.spinner_running is False diff --git a/tests/conftest.py b/tests/conftest.py index 27760ef231f..b5f08d51b63 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -14,3 +14,22 @@ def gpt35_model(): def gpt4_model(): """Common GPT-4 model fixture for tests requiring GPT-4.""" return Model("gpt-4") + + +# from pyinstrument import Profiler + +# @pytest.fixture(autouse=True, scope="session") +# def profile_suite(): +# profiler = Profiler() +# profiler.start() +# +# yield # The entire test suite runs here +# +# profiler.stop() +# +# # Save the interactive HTML report +# output_html_path = "pytest_profile.html" +# with open(output_html_path, "w", encoding="utf-8") as f: +# f.write(profiler.output_html()) +# +# print(f"\n[Pyinstrument] Flame graph saved to: {output_html_path}") diff --git a/tests/helpers/monorepo/test_repomap_workspace.py b/tests/helpers/monorepo/test_repomap_workspace.py index 0b14a760c45..5649889e535 100644 --- a/tests/helpers/monorepo/test_repomap_workspace.py +++ b/tests/helpers/monorepo/test_repomap_workspace.py @@ -17,7 +17,7 @@ def mock_workspace(tmp_path): # Project 1 p1_dir = workspace_root / "p1" / "main" p1_dir.mkdir(parents=True) - subprocess.run(["git", "init"], cwd=p1_dir, check=True) + subprocess.run(["git", "init", "-b", "main"], cwd=p1_dir, check=True) subprocess.run(["git", "config", "user.email", "test@test.com"], cwd=p1_dir, check=True) subprocess.run(["git", "config", "user.name", "Test"], cwd=p1_dir, check=True) (p1_dir / "file1.py").write_text("def func1(): pass") @@ -27,7 +27,7 @@ def mock_workspace(tmp_path): # Project 2 p2_dir = workspace_root / "p2" / "main" p2_dir.mkdir(parents=True) - subprocess.run(["git", "init"], cwd=p2_dir, check=True) + subprocess.run(["git", "init", "-b", "main"], cwd=p2_dir, check=True) subprocess.run(["git", "config", "user.email", "test@test.com"], cwd=p2_dir, check=True) subprocess.run(["git", "config", "user.name", "Test"], cwd=p2_dir, check=True) (p2_dir / "file2.py").write_text("def func2(): pass") diff --git a/tests/helpers/observations/test_observation_service.py b/tests/helpers/observations/test_observation_service.py index d1abdb63f09..3db9bb85664 100644 --- a/tests/helpers/observations/test_observation_service.py +++ b/tests/helpers/observations/test_observation_service.py @@ -67,6 +67,7 @@ async def test_compact_context_with_observations(): coder.context_compaction_summary_tokens = 100 coder.last_user_message = "Last user msg" coder.io = MagicMock() + coder.args = {} # Mock observation manager with some observations obs_manager = ObservationService.get_instance(coder) @@ -138,6 +139,7 @@ async def test_compact_context_with_observations_integration(): coder.context_compaction_summary_tokens = 100 coder.last_user_message = "Last user msg" coder.io = MagicMock() + coder.args = {} # Mock observation manager with some observations obs_manager = ObservationService.get_instance(coder) diff --git a/tests/mcp/test_keepalive_resilience.py b/tests/mcp/test_keepalive_resilience.py index 1bf90efe57a..0ea16931e0f 100644 --- a/tests/mcp/test_keepalive_resilience.py +++ b/tests/mcp/test_keepalive_resilience.py @@ -117,7 +117,7 @@ async def mock_sleep(duration): with patch("asyncio.sleep", side_effect=mock_sleep): await server.connect() # Yield control to the keepalive task multiple times - for _ in range(30): + for _ in range(200): await original_sleep(0) await server.disconnect() diff --git a/tests/mcp/test_manager_retry.py b/tests/mcp/test_manager_retry.py index 4cdf3de9c47..02d4fc35f23 100644 --- a/tests/mcp/test_manager_retry.py +++ b/tests/mcp/test_manager_retry.py @@ -434,7 +434,6 @@ async def test_connect_server_no_error_log_unnamed_server(mock_io): assert result is False mock_io.tool_error.assert_not_called() - assert mock_io.tool_warning.call_count == 2 assert unnamed_server not in manager._connected_servers