diff --git a/benchmark/README.md b/benchmark/README.md index 447267a0f90..7e8c23accd1 100644 --- a/benchmark/README.md +++ b/benchmark/README.md @@ -26,6 +26,13 @@ without any human review or supervision! The LLM could generate dangerous python that harms your system, like this: `import os; os.system("sudo rm -rf /")`. Running inside a docker container helps limit the damage that could be done. +The current harness (`benchmark.py`) organizes the exercises as "Cats": each +exercise carries `cat.yaml` metadata, enabling per-test selection and +deterministic subsets (e.g. `--hash-re '^0'`) for spot-checking models. The +canonical exercises source is +[ErichBSchulz/cecli-cats](https://github.com/ErichBSchulz/cecli-cats), a +Cats-formatted conversion of the polyglot benchmark. + ## Usage There are 3 main tasks involved in benchmarking: @@ -41,8 +48,8 @@ There are 3 main tasks involved in benchmarking: These steps only need to be done once. ``` -ORG=Aider-AI -REPO=aider +ORG=cecli-dev +REPO=cecli # Clone the main repo git clone https://github.com/$ORG/$REPO.git @@ -50,8 +57,9 @@ git clone https://github.com/$ORG/$REPO.git cd $REPO mkdir tmp.benchmarks -# Clone the repo with the exercises -git clone https://github.com/$ORG/polyglot-benchmark tmp.benchmarks/polyglot-benchmark +# Clone the canonical exercises source for the Cats harness +git clone https://github.com/ErichBSchulz/cecli-cats tmp.benchmarks/cecli-cats + # Build the docker container ./benchmark/docker_build.sh @@ -69,10 +77,10 @@ Launch the docker container and run the benchmark inside it: ./benchmark/docker.sh # Run the benchmark: -./benchmark/benchmark.py a-helpful-name-for-this-run --model gpt-3.5-turbo --edit-format whole --threads 10 --exercises-dir polyglot-benchmark +./benchmark/benchmark.py a-helpful-name-for-this-run --model gpt-3.5-turbo --edit-format whole --threads 10 --exercises-dir cecli-cats # Or with OpenRouter models (requires OPENROUTER_API_KEY environment variable): -./benchmark/benchmark.py openrouter-run --model openrouter/deepseek/deepseek-r1:free --edit-format whole --threads 10 --exercises-dir polyglot-benchmark +./benchmark/benchmark.py openrouter-run --model openrouter/deepseek/deepseek-r1:free --edit-format whole --threads 10 --exercises-dir cecli-cats ``` The above will create a folder @@ -156,6 +164,7 @@ Note the roadmap priorities: ## Limitations +- `benchmark_classic.py` is deprecated (it predates the async Coder API). - These scripts are not intended for use by typical `cecli` end users. - Some of the old (?deprecated) tools are written as `bash` scripts, so it will be hard to use them on Windows. diff --git a/cecli/__init__.py b/cecli/__init__.py index 66fb65c8e64..042a7acbe9f 100644 --- a/cecli/__init__.py +++ b/cecli/__init__.py @@ -1,6 +1,6 @@ from packaging import version -__version__ = "0.100.11.dev" +__version__ = "0.100.14.dev" safe_version = __version__ try: 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 2c9246b1424..bcb4d56e7fe 100644 --- a/cecli/mcp/manager.py +++ b/cecli/mcp/manager.py @@ -160,17 +160,46 @@ async def connect_server(self, name: str) -> bool: self._server_tools[server.name] = get_local_tool_schemas() return True - try: - session = await server.connect() - tools = await experimental_mcp_client.load_mcp_tools(session=session, format="openai") - self._server_tools[server.name] = tools - self._connected_servers.add(server) - self._log_verbose(f"Connected to MCP server: {name}") - return True - except (Exception, asyncio.CancelledError) as e: - if server.name != "unnamed-server": - self._log_error(f"Failed to connect to MCP server {name}: {e}") - return False + # Retry with exponential backoff for transient connection failures. + # Note: This also fixes a latent bug where asyncio.CancelledError was + # silently caught and treated as a connection failure. CancelledError is + # now re-raised to properly propagate cancellation. + # 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 if server.name != "unnamed-server" else 1 + delay = 1.0 + backoff = 2.0 + max_delay = 30.0 + + for attempt in range(1, max_retries + 1): + try: + session = await server.connect() + tools = await experimental_mcp_client.load_mcp_tools( + session=session, format="openai" + ) + self._server_tools[server.name] = tools + self._connected_servers.add(server) + self._log_verbose(f"Connected to MCP server: {name}") + return True + except asyncio.CancelledError: + raise + except Exception as e: + 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: + if server.name != "unnamed-server": + self._log_error( + f"Failed to connect to MCP server {name} " + f"after {max_retries} attempts: {e}" + ) + return False async def disconnect_server(self, name: str) -> bool: """ @@ -281,11 +310,10 @@ async def add_server_with_retry( success = await mcp_manager.add_server(server, connect=False) return (server, success) - for _attempt in range(max_retries): - success = await mcp_manager.add_server(server, connect=True) - if success: - return (server, True) - return (server, False) + # connect_server now has built-in retry logic, so we only need + # a single call here — no separate retry loop needed. + success = await mcp_manager.add_server(server, connect=True) + return (server, success) tasks = [] for server in servers: diff --git a/cecli/mcp/server.py b/cecli/mcp/server.py index bb92c1473dd..d130edb53a8 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).""" @@ -272,6 +314,8 @@ async def connect(self): server_info["client_info"]["token_endpoint"] = token_endpoint save_mcp_oauth_token(self.name, server_info) + + return session except Exception as e: logging.error(f"Error initializing {self.name}: {e}") await self.disconnect() 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.py b/tests/mcp/test_manager.py index 8c5ee5eb6b1..42b7391f7ef 100644 --- a/tests/mcp/test_manager.py +++ b/tests/mcp/test_manager.py @@ -149,11 +149,15 @@ async def test_connect_server_failure(self, mock_server, mock_io): manager = McpServerManager(servers=[mock_server], io=mock_io) mock_server.connect.side_effect = Exception("Connection failed") - result = await manager.connect_server("test-server") + with patch("asyncio.sleep"): + result = await manager.connect_server("test-server") assert result is False - mock_server.connect.assert_called_once() + assert mock_server.connect.call_count == 3 # 1 initial + 2 retries + assert mock_io.tool_warning.call_count == 2 # warnings for attempts 1 and 2 mock_io.tool_error.assert_called_once() + error_msg = mock_io.tool_error.call_args[0][0] + assert "after 3 attempts" in error_msg assert mock_server not in manager._connected_servers @pytest.mark.asyncio diff --git a/tests/mcp/test_manager_retry.py b/tests/mcp/test_manager_retry.py new file mode 100644 index 00000000000..02d4fc35f23 --- /dev/null +++ b/tests/mcp/test_manager_retry.py @@ -0,0 +1,826 @@ +"""Comprehensive retry logic tests for McpServerManager.connect_server. + +This file implements all 28 test cases (TC-001 through TC-028) from the +MCP retry logic plan in .cecli.plans.md Sections 10 and 11. + +Test categories: + - Core retry behavior (TC-001, TC-002, TC-003, TC-014) + - Edge cases - cancellation (TC-007, TC-008) + - Edge cases - no retry (TC-004, TC-005, TC-006, TC-013) + - Timing tests (TC-009, TC-010) + - Logging tests (TC-011, TC-012) + - Integration tests (TC-015 through TC-021) + - Regression tests (TC-022 through TC-027) + - Special case - tool loading failure (TC-028) +""" + +import asyncio +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from cecli.commands.core import SwitchCoderSignal +from cecli.mcp.manager import McpServerManager +from cecli.mcp.server import LocalServer, McpServer + +# --------------------------------------------------------------------------- +# Fixtures (mirrors tests/mcp/test_manager.py for consistency) +# --------------------------------------------------------------------------- + + +@pytest.fixture +def mock_io(): + """Mock IO object for capturing log calls.""" + io = MagicMock() + io.tool_output = MagicMock() + io.tool_error = MagicMock() + io.tool_warning = MagicMock() + return io + + +@pytest.fixture +def mock_server(): + """Mock McpServer named 'test-server' with async connect/disconnect.""" + server = MagicMock(spec=McpServer) + server.name = "test-server" + server.config = {"name": "test-server", "enabled": True} + server.connect = AsyncMock() + server.disconnect = AsyncMock() + server.is_connected = False + return server + + +@pytest.fixture +def mock_local_server(): + """Mock LocalServer named 'Local'.""" + server = MagicMock(spec=LocalServer) + server.name = "Local" + server.config = {"name": "Local", "enabled": True} + server.connect = AsyncMock() + server.disconnect = AsyncMock() + server.is_connected = False + return server + + +@pytest.fixture +def mock_tools(): + """Mock tool schemas returned by load_mcp_tools.""" + return [ + { + "function": { + "name": "test_tool", + "description": "A test tool", + "parameters": {}, + } + } + ] + + +@pytest.fixture +def mock_session(): + """Mock session object returned by server.connect().""" + return MagicMock() + + +# --------------------------------------------------------------------------- +# Helper to create a mock coder for integration tests +# --------------------------------------------------------------------------- + + +def _make_mock_coder(mock_server, mock_io, connected_servers=None): + """Create a mock coder with mcp_manager for integration tests. + + Args: + mock_server: The mock server to include in the manager. + mock_io: Mock IO object. + connected_servers: Optional set of already-connected servers. + + Returns: + A MagicMock coder with mcp_manager, coroutines, interrupt_event. + """ + coder = MagicMock() + coder.io = mock_io + coder.edit_format = "agent" + coder.interrupt_event = MagicMock() + coder.interrupt_event.clear = MagicMock() + coder.interrupt_event.is_set = MagicMock(return_value=False) + + # Create mcp_manager + coder.mcp_manager = MagicMock() + coder.mcp_manager.servers = [mock_server] + coder.mcp_manager.connected_servers = connected_servers or [] + coder.mcp_manager.get_server = MagicMock(return_value=mock_server) + coder.mcp_manager.connect_server = AsyncMock() + + # Create coroutines with interruptible that passes through + coder.coroutines = MagicMock() + + async def _passthrough_interruptible(coro, event): + """Pass-through interruptible that just awaits the coroutine.""" + return await coro, False + + coder.coroutines.interruptible = _passthrough_interruptible + + # registered_servers for update_server_registration + coder.registered_servers = {"included": set(), "excluded": set()} + + return coder + + +# --------------------------------------------------------------------------- +# TC-001: connect_server retries on first failure, succeeds on second attempt +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_retries_first_failure_succeeds_second( + mock_server, mock_io, mock_tools, mock_session +): + """TC-001: connect_server retries after first connection failure and succeeds on second attempt.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = [Exception("Connection failed"), mock_session] + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.return_value = mock_tools + with patch("asyncio.sleep") as mock_sleep: + result = await manager.connect_server("test-server") + + assert result is True + assert mock_server.connect.call_count == 2 + assert mock_sleep.call_count == 1 + assert mock_sleep.call_args[0][0] == 1.0 + assert mock_io.tool_warning.call_count == 1 + assert mock_server in manager._connected_servers + assert manager._server_tools["test-server"] == mock_tools + + +# --------------------------------------------------------------------------- +# TC-002: connect_server retries on first and second failure, succeeds on third +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_retries_two_failures_succeeds_third( + mock_server, mock_io, mock_tools, mock_session +): + """TC-002: connect_server retries after two failures and succeeds on third attempt.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = [ + Exception("Fail 1"), + Exception("Fail 2"), + mock_session, + ] + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.return_value = mock_tools + with patch("asyncio.sleep") as mock_sleep: + result = await manager.connect_server("test-server") + + assert result is True + assert mock_server.connect.call_count == 3 + assert mock_sleep.call_count == 2 + assert mock_sleep.call_args_list[0][0][0] == 1.0 + assert mock_sleep.call_args_list[1][0][0] == 2.0 + assert mock_io.tool_warning.call_count == 2 + assert mock_server in manager._connected_servers + + +# --------------------------------------------------------------------------- +# TC-003: connect_server returns False after all 3 retries exhausted +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_fails_after_all_retries(mock_server, mock_io): + """TC-003: connect_server fails after all 3 retry attempts are exhausted.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = Exception("Connection failed") + + with patch("asyncio.sleep"): + result = await manager.connect_server("test-server") + + assert result is False + assert mock_server.connect.call_count == 3 + assert mock_io.tool_warning.call_count == 2 + mock_io.tool_error.assert_called_once() + error_msg = mock_io.tool_error.call_args[0][0] + assert "after 3 attempts" in error_msg + assert mock_server not in manager._connected_servers + assert "test-server" not in manager._server_tools + + +# --------------------------------------------------------------------------- +# TC-004: connect_server does not retry for LocalServer +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_no_retry_local_server(mock_local_server): + """TC-004: connect_server connects LocalServer on first attempt without retry.""" + manager = McpServerManager(servers=[mock_local_server]) + + with patch("cecli.mcp.manager.get_local_tool_schemas") as mock_get_schemas: + mock_get_schemas.return_value = [{"name": "local_tool"}] + result = await manager.connect_server("Local") + + assert result is True + assert mock_local_server.connect.call_count == 1 + assert mock_local_server in manager._connected_servers + assert manager._server_tools["Local"] == [{"name": "local_tool"}] + + +# --------------------------------------------------------------------------- +# TC-005: connect_server does not retry when server not found +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_no_retry_not_found(mock_io): + """TC-005: connect_server returns False immediately for non-existent server.""" + manager = McpServerManager(servers=[], io=mock_io) + + result = await manager.connect_server("nonexistent-server") + + assert result is False + mock_io.tool_warning.assert_called_once() + warning_msg = mock_io.tool_warning.call_args[0][0] + assert "not found" in warning_msg + assert len(manager._connected_servers) == 0 + + +# --------------------------------------------------------------------------- +# TC-006: connect_server does not retry when already connected +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_no_retry_already_connected(mock_server, mock_io): + """TC-006: connect_server returns True immediately for already-connected server.""" + manager = McpServerManager(servers=[mock_server], io=mock_io, verbose=True) + manager._connected_servers.add(mock_server) + + result = await manager.connect_server("test-server") + + assert result is True + mock_server.connect.assert_not_called() + mock_io.tool_output.assert_called_once() + output_msg = mock_io.tool_output.call_args[0][0] + assert "already connected" in output_msg + + +# --------------------------------------------------------------------------- +# TC-007: connect_server propagates asyncio.CancelledError during retry delay +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_propagates_cancelled_error_during_retry(mock_server, mock_io): + """TC-007: connect_server re-raises CancelledError when interrupted during retry backoff.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = Exception("Connection failed") + + with patch("asyncio.sleep") as mock_sleep: + mock_sleep.side_effect = asyncio.CancelledError() + with pytest.raises(asyncio.CancelledError): + await manager.connect_server("test-server") + + assert mock_server.connect.call_count == 1 + assert mock_sleep.call_count == 1 + mock_io.tool_error.assert_not_called() + + +# --------------------------------------------------------------------------- +# TC-008: connect_server propagates asyncio.CancelledError during server.connect() +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_propagates_cancelled_error_during_connect(mock_server, mock_io): + """TC-008: connect_server re-raises CancelledError when server.connect() raises it.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = asyncio.CancelledError() + + with pytest.raises(asyncio.CancelledError): + await manager.connect_server("test-server") + + assert mock_server.connect.call_count == 1 + mock_io.tool_error.assert_not_called() + + +# --------------------------------------------------------------------------- +# TC-009: connect_server uses exponential backoff timing (1s, 2s, 4s) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_exponential_backoff_timing(mock_server, mock_io): + """TC-009: connect_server applies exponential backoff delays of 1s, 2s between retries.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = Exception("Connection failed") + + sleep_delays = [] + + async def _capture_sleep(delay): + sleep_delays.append(delay) + + with patch("asyncio.sleep", side_effect=_capture_sleep): + result = await manager.connect_server("test-server") + + assert result is False + assert len(sleep_delays) == 2 + assert sleep_delays[0] == 1.0 + assert sleep_delays[1] == 2.0 + + +# --------------------------------------------------------------------------- +# TC-010: connect_server backoff is capped at max_delay of 30 seconds +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_backoff_capped_at_max_delay(mock_server, mock_io): + """TC-010: connect_server backoff delay does not exceed 30 seconds maximum. + + With default max_retries=3, delays are 1.0 and 2.0 (both below 30s cap). + We verify the min(delay * backoff, max_delay) logic by checking that + all recorded delays are <= 30.0. + """ + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = Exception("Connection failed") + + sleep_delays = [] + + async def _capture_sleep(delay): + sleep_delays.append(delay) + + with patch("asyncio.sleep", side_effect=_capture_sleep): + result = await manager.connect_server("test-server") + + assert result is False + assert len(sleep_delays) == 2 + for delay in sleep_delays: + assert delay <= 30.0 + assert sleep_delays[0] == 1.0 + assert sleep_delays[1] == 2.0 + + +# --------------------------------------------------------------------------- +# TC-011: connect_server logs warning on each failed attempt (except final) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_logs_warning_on_failed_attempts(mock_server, mock_io): + """TC-011: connect_server calls _log_warning for each non-final failed attempt.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = Exception("Connection failed") + + with patch("asyncio.sleep"): + await manager.connect_server("test-server") + + assert mock_io.tool_warning.call_count == 2 + + warning1 = mock_io.tool_warning.call_args_list[0][0][0] + assert "attempt 1 failed" in warning1 + assert "retrying in 1.0s" in warning1 + assert "Connection failed" in warning1 + + warning2 = mock_io.tool_warning.call_args_list[1][0][0] + assert "attempt 2 failed" in warning2 + assert "retrying in 2.0s" in warning2 + assert "Connection failed" in warning2 + + +# --------------------------------------------------------------------------- +# TC-012: connect_server logs error on final failure with attempt count +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_logs_error_on_final_failure(mock_server, mock_io): + """TC-012: connect_server calls _log_error after all retries exhausted.""" + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = Exception("Connection failed") + + with patch("asyncio.sleep"): + await manager.connect_server("test-server") + + mock_io.tool_error.assert_called_once() + error_msg = mock_io.tool_error.call_args[0][0] + assert "Failed to connect to MCP server" in error_msg + assert "after 3 attempts" in error_msg + assert "Connection failed" in error_msg + + +# --------------------------------------------------------------------------- +# TC-013: connect_server does not log error for "unnamed-server" on final failure +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_no_error_log_unnamed_server(mock_io): + """TC-013: connect_server suppresses error logging for servers named 'unnamed-server'.""" + unnamed_server = MagicMock(spec=McpServer) + unnamed_server.name = "unnamed-server" + unnamed_server.config = {"name": "unnamed-server", "enabled": True} + unnamed_server.connect = AsyncMock(side_effect=Exception("Connection failed")) + unnamed_server.disconnect = AsyncMock() + unnamed_server.is_connected = False + + manager = McpServerManager(servers=[unnamed_server], io=mock_io) + + with patch("asyncio.sleep"): + result = await manager.connect_server("unnamed-server") + + assert result is False + mock_io.tool_error.assert_not_called() + assert unnamed_server not in manager._connected_servers + + +# --------------------------------------------------------------------------- +# TC-014: connect_server succeeds on first attempt (no retry needed) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_succeeds_first_attempt(mock_server, mock_tools, mock_session): + """TC-014: connect_server connects successfully on first attempt without any retries.""" + manager = McpServerManager(servers=[mock_server]) + mock_server.connect.return_value = mock_session + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.return_value = mock_tools + result = await manager.connect_server("test-server") + + assert result is True + assert mock_server.connect.call_count == 1 + assert mock_server in manager._connected_servers + assert manager._server_tools["test-server"] == mock_tools + mock_load_tools.assert_called_once_with(session=mock_session, format="openai") + + +# --------------------------------------------------------------------------- +# TC-015: from_servers simplification - add_server_with_retry calls connect once +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_from_servers_calls_connect_once(mock_server, mock_io, mock_tools, mock_session): + """TC-015: from_servers add_server_with_retry no longer has its own retry loop.""" + mock_server.connect.return_value = mock_session + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.return_value = mock_tools + manager = await McpServerManager.from_servers( + servers=[mock_server], io=mock_io, verbose=True + ) + + assert isinstance(manager, McpServerManager) + assert manager._servers == [mock_server] + assert mock_server in manager._connected_servers + assert mock_server.connect.call_count == 1 + mock_load_tools.assert_called_once() + + +# --------------------------------------------------------------------------- +# TC-016: from_servers shows warning for failed server after connect_server retries +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_from_servers_warning_after_retries(mock_server, mock_io): + """TC-016: from_servers displays warning when server fails to connect after retries.""" + mock_server.connect.side_effect = Exception("Connection failed") + + with patch("asyncio.sleep"): + manager = await McpServerManager.from_servers(servers=[mock_server], io=mock_io) + + assert isinstance(manager, McpServerManager) + assert mock_server not in manager._connected_servers + + warning_messages = [call[0][0] for call in mock_io.tool_warning.call_args_list] + found_init_warning = any("MCP tool initialization failed" in msg for msg in warning_messages) + assert found_init_warning + + assert mock_server.connect.call_count == 3 + + +# --------------------------------------------------------------------------- +# TC-017: /load-mcp command benefits from retry logic +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_load_mcp_command_benefits_from_retry(mock_server, mock_io, mock_tools, mock_session): + """TC-017: LoadMcpCommand.execute uses connect_server which now has built-in retry.""" + from cecli.commands.load_mcp import LoadMcpCommand + + coder = _make_mock_coder(mock_server, mock_io) + + # connect_server fails first, succeeds second (simulating retry inside) + async def _connect_with_retry(name): + # Simulate the retry happening inside connect_server + mock_server.connect.side_effect = [Exception("Fail"), mock_session] + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load: + mock_load.return_value = mock_tools + with patch("asyncio.sleep"): + manager = McpServerManager(servers=[mock_server], io=mock_io) + return await manager.connect_server(name) + + coder.mcp_manager.connect_server = _connect_with_retry + coder.mcp_manager.get_server = MagicMock(return_value=mock_server) + coder.mcp_manager.connected_servers = [] + + with patch("cecli.commands.load_mcp.iter_all_coders", return_value=[coder]): + with patch("cecli.commands.load_mcp.update_server_registration"): + with pytest.raises(SwitchCoderSignal): + await LoadMcpCommand.execute(mock_io, coder, "test-server") + + +# --------------------------------------------------------------------------- +# TC-018: /load-mcp command reports failure after retries exhausted +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_load_mcp_command_reports_failure_after_retries(mock_server, mock_io): + """TC-018: LoadMcpCommand.execute reports 'Unable to load server' after retries exhausted.""" + from cecli.commands.load_mcp import LoadMcpCommand + + coder = _make_mock_coder(mock_server, mock_io) + coder.mcp_manager.connect_server = AsyncMock(return_value=False) + coder.mcp_manager.get_server = MagicMock(return_value=mock_server) + coder.mcp_manager.connected_servers = [] + + with patch("cecli.commands.load_mcp.iter_all_coders", return_value=[coder]): + with patch("cecli.commands.load_mcp.update_server_registration"): + with pytest.raises(SwitchCoderSignal): + await LoadMcpCommand.execute(mock_io, coder, "test-server") + + # Check that the results were output - "Unable to load server" should be in output + output_calls = [str(call) for call in mock_io.tool_output.call_args_list] + all_output = " ".join(output_calls) + assert "Unable to load server: test-server" in all_output + + +# --------------------------------------------------------------------------- +# TC-019: /load-mcp interruptible wrapper propagates cancellation during retry +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_load_mcp_propagates_cancellation_during_retry(mock_server, mock_io): + """TC-019: LoadMcpCommand interruptible wrapper correctly propagates CancelledError.""" + from cecli.commands.load_mcp import LoadMcpCommand + + coder = _make_mock_coder(mock_server, mock_io) + + async def _connect_raises_cancelled(name): + raise asyncio.CancelledError() + + coder.mcp_manager.connect_server = _connect_raises_cancelled + coder.mcp_manager.get_server = MagicMock(return_value=mock_server) + coder.mcp_manager.connected_servers = [] + + # Make interruptible propagate the CancelledError + async def _propagate_interruptible(coro, event): + return await coro, False + + coder.coroutines.interruptible = _propagate_interruptible + + with patch("cecli.commands.load_mcp.iter_all_coders", return_value=[coder]): + with patch("cecli.commands.load_mcp.update_server_registration"): + with pytest.raises(asyncio.CancelledError): + await LoadMcpCommand.execute(mock_io, coder, "test-server") + + +# --------------------------------------------------------------------------- +# TC-020: /load-session command benefits from retry logic +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_load_session_benefits_from_retry(mock_server, mock_io, mock_tools, mock_session): + """TC-020: Session loading uses connect_server which now retries on transient failures. + + This test verifies that connect_server is called with the correct server name + and that retry logic is exercised during session loading. + """ + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = [Exception("Fail"), mock_session] + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.return_value = mock_tools + with patch("asyncio.sleep"): + # Simulate what /load-session does: call connect_server for the MCP server + result = await manager.connect_server("test-server") + + assert result is True + assert mock_server.connect.call_count == 2 + assert mock_server in manager._connected_servers + assert manager._server_tools["test-server"] == mock_tools + + +# --------------------------------------------------------------------------- +# TC-021: Resource Manager tool benefits from retry logic +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_resource_manager_benefits_from_retry(mock_server, mock_io, mock_tools, mock_session): + """TC-021: ResourceManager _load_mcp uses connect_server which now retries.""" + from cecli.tools.resource_manager import Tool as ResourceManagerTool + + coder = _make_mock_coder(mock_server, mock_io) + + # Simulate connect_server with retry (fail first, succeed second) + async def _connect_with_retry(name): + mock_server.connect.side_effect = [Exception("Fail"), mock_session] + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load: + mock_load.return_value = mock_tools + with patch("asyncio.sleep"): + manager = McpServerManager(servers=[mock_server], io=mock_io) + return await manager.connect_server(name) + + coder.mcp_manager.connect_server = _connect_with_retry + coder.mcp_manager.get_server = MagicMock(return_value=mock_server) + coder.mcp_manager.connected_servers = [] + + # Mock the context block check + coder.agent_config = {"include_context_blocks": {"servers"}, "exclude_context_blocks": set()} + + with patch("cecli.tools.resource_manager.iter_all_coders", return_value=[coder]): + with patch("cecli.tools.resource_manager.update_server_registration"): + result = await ResourceManagerTool._load_mcp(coder, "test-server") + + assert "Loaded server: test-server" in result + + +# --------------------------------------------------------------------------- +# TC-022: Regression - existing connect_server_success test still passes +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_regression_connect_server_success(mock_server, mock_tools, mock_session): + """TC-022: Regression - connect_server succeeds on first attempt (existing behavior).""" + manager = McpServerManager(servers=[mock_server]) + mock_server.connect.return_value = mock_session + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.return_value = mock_tools + result = await manager.connect_server("test-server") + + assert result is True + assert mock_server.connect.call_count == 1 + assert mock_server in manager._connected_servers + assert manager._server_tools["test-server"] == mock_tools + + +# --------------------------------------------------------------------------- +# TC-023: Regression - existing connect_server_failure test still passes +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_regression_connect_server_failure(mock_server, mock_io): + """TC-023: Regression - connect_server failure now retries 3 times (updated behavior). + + Note: The original test_connect_server_failure in test_manager.py asserted + connect() called once. With retry logic, connect() is called 3 times. + The existing test in test_manager.py has been updated to reflect this. + This test replicates those assertions here for regression coverage. + """ + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.side_effect = Exception("Connection failed") + + with patch("asyncio.sleep"): + result = await manager.connect_server("test-server") + + assert result is False + assert mock_server.connect.call_count == 3 + assert mock_io.tool_warning.call_count == 2 + mock_io.tool_error.assert_called_once() + error_msg = mock_io.tool_error.call_args[0][0] + assert "after 3 attempts" in error_msg + assert mock_server not in manager._connected_servers + + +# --------------------------------------------------------------------------- +# TC-024: Regression - existing connect_server_not_found test still passes +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_regression_connect_server_not_found(mock_io): + """TC-024: Regression - connect_server returns False immediately for not-found server.""" + manager = McpServerManager(servers=[], io=mock_io) + + result = await manager.connect_server("nonexistent-server") + + assert result is False + mock_io.tool_warning.assert_called_once() + + +# --------------------------------------------------------------------------- +# TC-025: Regression - existing connect_server_already_connected test still passes +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_regression_connect_server_already_connected(mock_server, mock_io): + """TC-025: Regression - connect_server returns True immediately for already-connected server.""" + manager = McpServerManager(servers=[mock_server], io=mock_io, verbose=True) + manager._connected_servers.add(mock_server) + + result = await manager.connect_server("test-server") + + assert result is True + mock_server.connect.assert_not_called() + mock_io.tool_output.assert_called_once() + + +# --------------------------------------------------------------------------- +# TC-026: Regression - existing connect_server_local_server test still passes +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_regression_connect_server_local_server(mock_local_server): + """TC-026: Regression - connect_server connects LocalServer without retry.""" + manager = McpServerManager(servers=[mock_local_server]) + + with patch("cecli.mcp.manager.get_local_tool_schemas") as mock_get_schemas: + mock_get_schemas.return_value = [{"name": "local_tool"}] + result = await manager.connect_server("Local") + + assert result is True + assert mock_local_server.connect.call_count == 1 + assert mock_local_server in manager._connected_servers + assert manager._server_tools["Local"] == [{"name": "local_tool"}] + + +# --------------------------------------------------------------------------- +# TC-027: Regression - existing from_servers_creates_manager test still passes +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_regression_from_servers_creates_manager( + mock_server, mock_io, mock_tools, mock_session +): + """TC-027: Regression - from_servers creates manager with simplified add_server_with_retry.""" + mock_server.connect.return_value = mock_session + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.return_value = mock_tools + manager = await McpServerManager.from_servers( + servers=[mock_server], io=mock_io, verbose=True + ) + + assert isinstance(manager, McpServerManager) + assert manager._servers == [mock_server] + assert mock_server in manager._connected_servers + assert mock_server.connect.call_count == 1 + mock_load_tools.assert_called_once() + + +# --------------------------------------------------------------------------- +# TC-028: connect_server retries on tool-loading failure +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_connect_server_retries_on_tool_loading_failure( + mock_server, mock_io, mock_tools, mock_session +): + """TC-028: connect_server retries when connect() succeeds but load_mcp_tools() fails. + + In this scenario: + - server.connect() succeeds (returns session) + - load_mcp_tools() fails on first call, succeeds on second + - The retry loop catches the exception from load_mcp_tools and retries + - server.connect() is called again on retry (it returns the same session) + - load_mcp_tools() is called twice total + """ + manager = McpServerManager(servers=[mock_server], io=mock_io) + mock_server.connect.return_value = mock_session + + with patch("litellm.experimental_mcp_client.load_mcp_tools") as mock_load_tools: + mock_load_tools.side_effect = [Exception("Tool load failed"), mock_tools] + with patch("asyncio.sleep") as mock_sleep: + result = await manager.connect_server("test-server") + + assert result is True + # connect() is called for each attempt (2 attempts: first fails at load_mcp_tools, second succeeds) + assert mock_server.connect.call_count == 2 + # load_mcp_tools() called twice: first fails, second succeeds + assert mock_load_tools.call_count == 2 + # sleep called once (between attempt 1 and 2) + assert mock_sleep.call_count == 1 + assert mock_sleep.call_args[0][0] == 1.0 + # Warning logged for the first failed attempt + assert mock_io.tool_warning.call_count == 1 + # Server connected and tools loaded + assert mock_server in manager._connected_servers + assert manager._server_tools["test-server"] == mock_tools