From b559a910977e2e152b61d1f842316406b2d1e736 Mon Sep 17 00:00:00 2001 From: Petr Date: Thu, 5 Mar 2026 15:05:17 +0100 Subject: [PATCH 1/3] MCP performance: persistent HTTP server + single-session calls - Reorder detect_mcp_server_command(): local first, python -m second, uvx last. Remove @latest from uvx to avoid PyPI check (~25s penalty). - Merge validate + call into single MCP session via new validate_and_call_tool() method, eliminating double subprocess spawn. - Add persistent HTTP server transport (McpServerManager) that keeps one keboola-mcp-server running with streamable-http transport. Per-request project credentials via X-Storage-Token/X-Storage-API-URL headers allow one server to serve all projects. - KBAGENT_MCP_TRANSPORT env var controls mode: "http" (default) or "stdio". - Update doctor command to show transport mode and server status. - Add test_mcp_transport.py with 20 tests for server manager. - Update existing tests for new detection order and unified call flow. Expected improvement: ~50s -> ~3-6s per tool call (local/cached), ~2-3s for subsequent calls via persistent HTTP server. Co-Authored-By: Claude Opus 4.6 --- src/keboola_agent_cli/commands/tool.py | 41 +- src/keboola_agent_cli/constants.py | 9 + src/keboola_agent_cli/services/mcp_service.py | 623 +++++++++++++++++- .../services/mcp_transport.py | 215 ++++++ tests/conftest.py | 6 + tests/test_cli.py | 31 +- tests/test_mcp_service.py | 75 ++- tests/test_mcp_transport.py | 265 ++++++++ 8 files changed, 1170 insertions(+), 95 deletions(-) create mode 100644 src/keboola_agent_cli/services/mcp_transport.py create mode 100644 tests/test_mcp_transport.py diff --git a/src/keboola_agent_cli/commands/tool.py b/src/keboola_agent_cli/commands/tool.py index d24fade9..db106fbc 100644 --- a/src/keboola_agent_cli/commands/tool.py +++ b/src/keboola_agent_cli/commands/tool.py @@ -207,41 +207,26 @@ def tool_call( ) raise typer.Exit(code=2) from None - # Validate required parameters before calling tool across all projects + # Validate + call in a single MCP session (eliminates double subprocess spawn) try: - missing, known_tools = service.validate_tool_input( - tool_name=tool_name, - tool_input=parsed_input, - aliases=[project] if project else None, - branch_id=branch_str, - ) - except ConfigError as exc: - formatter.error(message=exc.message, error_code="CONFIG_ERROR") - raise typer.Exit(code=5) from None - - if missing: - params_str = ", ".join(missing) - example_json = json.dumps({p: "..." for p in missing}) - formatter.error( - message=( - f"Missing required parameter(s) for '{tool_name}': {params_str}. " - f"Use: kbagent tool call {tool_name} --input '{example_json}'" - ), - error_code="MISSING_PARAMETER", - ) - raise typer.Exit(code=2) from None - - try: - result = service.call_tool( + result = service.validate_and_call_tool( tool_name=tool_name, tool_input=parsed_input, alias=project, branch_id=branch_str, - _known_tools=known_tools, ) except ConfigError as exc: - formatter.error(message=exc.message, error_code="CONFIG_ERROR") - raise typer.Exit(code=5) from None + # ConfigError covers: unknown tool, missing params, config issues + error_code = "CONFIG_ERROR" + exit_code = 5 + if "Missing required parameter" in exc.message: + error_code = "MISSING_PARAMETER" + exit_code = 2 + elif "Unknown MCP tool" in exc.message: + error_code = "CONFIG_ERROR" + exit_code = 5 + formatter.error(message=exc.message, error_code=error_code) + raise typer.Exit(code=exit_code) from None if formatter.json_mode: formatter.output(result) diff --git a/src/keboola_agent_cli/constants.py b/src/keboola_agent_cli/constants.py index a808ae1b..5e55b5e0 100644 --- a/src/keboola_agent_cli/constants.py +++ b/src/keboola_agent_cli/constants.py @@ -41,6 +41,15 @@ # 0 = unlimited (all projects run in parallel); set KBAGENT_MCP_MAX_SESSIONS to throttle DEFAULT_MCP_MAX_SESSIONS: int = 0 +# --- MCP HTTP Transport --- +# Transport mode: "http" (persistent server) or "stdio" (subprocess per call) +ENV_MCP_TRANSPORT: str = "KBAGENT_MCP_TRANSPORT" +DEFAULT_MCP_TRANSPORT: str = "http" +# Timeout for the persistent MCP server to start and be healthy +MCP_SERVER_STARTUP_TIMEOUT: float = 15.0 +# Timeout for health check requests to persistent MCP server +MCP_SERVER_HEALTH_TIMEOUT: float = 2.0 + # --- Storage Job Polling --- STORAGE_JOB_POLL_INTERVAL: float = 1.0 # seconds between polls STORAGE_JOB_MAX_WAIT: float = 60.0 # max seconds to wait for a storage job diff --git a/src/keboola_agent_cli/services/mcp_service.py b/src/keboola_agent_cli/services/mcp_service.py index d1340bd8..bb034726 100644 --- a/src/keboola_agent_cli/services/mcp_service.py +++ b/src/keboola_agent_cli/services/mcp_service.py @@ -1,7 +1,12 @@ -"""MCP integration service - wraps keboola-mcp-server as subprocess. +"""MCP integration service - wraps keboola-mcp-server. Provides multi-project tool listing and execution via MCP protocol. Read tools run across ALL projects in parallel; write tools target a single project. + +Supports two transport modes: +- HTTP (default): Persistent server with per-request credentials via headers. + One server serves all projects. Fastest for repeated calls. +- stdio: Subprocess per call. Fallback when HTTP transport is unavailable. """ import asyncio @@ -15,14 +20,17 @@ from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import stdio_client +from mcp.client.streamable_http import streamablehttp_client from ..constants import ( DEFAULT_MCP_INIT_TIMEOUT, DEFAULT_MCP_MAX_SESSIONS, DEFAULT_MCP_TOOL_TIMEOUT, + DEFAULT_MCP_TRANSPORT, ENV_MCP_INIT_TIMEOUT, ENV_MCP_MAX_SESSIONS, ENV_MCP_TOOL_TIMEOUT, + ENV_MCP_TRANSPORT, ) from ..errors import ConfigError from ..models import ProjectConfig @@ -82,20 +90,33 @@ def _is_write_tool(tool_name: str) -> bool: def detect_mcp_server_command() -> list[str] | None: """Detect the best way to run keboola-mcp-server. - Checks in order: - 1. uvx keboola_mcp_server@latest (if uvx is available -- always latest version) - 2. keboola_mcp_server (if installed as standalone command) - 3. python -m keboola_mcp_server (last resort) + Checks in order of speed (fastest first): + 1. keboola_mcp_server (local install, ~3s startup) + 2. python -m keboola_mcp_server (installed in current env) + 3. uvx keboola_mcp_server (cached version, ~4.5s startup) + + Note: We intentionally do NOT use @latest with uvx because it forces + a PyPI check on every invocation (~25s penalty). The cached version + is used instead. Users can update manually with: uvx upgrade keboola_mcp_server Returns: List of command parts, or None if no method is available. """ - if shutil.which("uvx"): - return ["uvx", "--prerelease=allow", "keboola_mcp_server@latest"] + # 1. Local install (fastest: ~3s) if shutil.which("keboola_mcp_server"): return ["keboola_mcp_server"] + # 2. python -m (if installed in current env) if shutil.which("python"): - return ["python", "-m", "keboola_mcp_server"] + result = subprocess.run( + ["python", "-c", "import keboola_mcp_server"], + capture_output=True, + timeout=5, + ) + if result.returncode == 0: + return ["python", "-m", "keboola_mcp_server"] + # 3. uvx WITHOUT @latest (uses cached version: ~4.5s vs 25s) + if shutil.which("uvx"): + return ["uvx", "--prerelease=allow", "keboola_mcp_server"] return None @@ -258,6 +279,339 @@ async def _connect_and_call_tool( await exit_stack.aclose() +async def _connect_validate_and_call( + project: ProjectConfig, + tool_name: str, + tool_input: dict[str, Any], + branch_id: str | None = None, +) -> dict[str, Any]: + """Open ONE MCP session: validate tool name + schema, then call the tool. + + Eliminates the need for a separate validate_tool_input() + call_tool() + sequence, saving one full subprocess spawn. + + Args: + project: Project config. + tool_name: Name of the MCP tool to call. + tool_input: Input arguments for the tool. + branch_id: Optional development branch ID. + + Returns: + Dict with "content", "isError", and "tool_schema" keys. + + Raises: + ConfigError: If tool_name is not found in the available tool list. + ConfigError: If required parameters are missing (not auto-expandable). + """ + exit_stack = AsyncExitStack() + + try: + session = await _open_session(project, exit_stack, branch_id=branch_id) + + # Step 1: list_tools to validate tool name and get schema + response = await asyncio.wait_for( + session.list_tools(), timeout=_get_tool_timeout() + ) + + known_tools = {t.name for t in response.tools} + if tool_name not in known_tools: + raise ConfigError( + f"Unknown MCP tool '{tool_name}'. " + f"Use 'kbagent tool list' to see available tools." + ) + + # Step 2: validate required params + schema: dict[str, Any] = {} + for tool in response.tools: + if tool.name == tool_name: + schema = tool.inputSchema if tool.inputSchema else {} + break + + required = schema.get("required", []) + missing = [param for param in required if param not in tool_input] + + # Exclude auto-expandable params from missing list + expand_config = AUTO_EXPAND_TOOLS.get(tool_name) + if expand_config: + auto_param = expand_config["param"] + missing = [p for p in missing if p != auto_param] + + if missing: + import json as _json + + params_str = ", ".join(missing) + example_json = _json.dumps({p: "..." for p in missing}) + raise ConfigError( + f"Missing required parameter(s) for '{tool_name}': {params_str}. " + f"Use: kbagent tool call {tool_name} --input '{example_json}'" + ) + + # Step 3: call the tool in the same session + result = await asyncio.wait_for( + session.call_tool(tool_name, tool_input), + timeout=_get_tool_timeout(), + ) + + return { + "content": _parse_content(result), + "isError": bool(result.isError), + } + finally: + await exit_stack.aclose() + + +def _get_transport_mode() -> str: + """Get configured MCP transport mode ('http' or 'stdio').""" + return os.environ.get(ENV_MCP_TRANSPORT, DEFAULT_MCP_TRANSPORT) + + +def _build_http_headers( + project: ProjectConfig, + branch_id: str | None = None, +) -> dict[str, str]: + """Build HTTP headers for per-request project credentials.""" + headers = { + "X-Storage-Token": project.token, + "X-Storage-API-URL": project.stack_url, + } + if branch_id: + headers["X-Branch-ID"] = branch_id + return headers + + +async def _http_list_tools( + base_url: str, + project: ProjectConfig, + branch_id: str | None = None, +) -> list[dict[str, Any]]: + """List tools via HTTP transport (persistent server). + + Args: + base_url: Base URL of the persistent MCP server. + project: Project config for authentication headers. + branch_id: Optional development branch ID. + + Returns: + List of tool dicts with name, description, inputSchema. + """ + headers = _build_http_headers(project, branch_id) + url = f"{base_url}/mcp" + + async with streamablehttp_client(url=url, headers=headers) as ( + read_stream, + write_stream, + _, + ): + session = ClientSession(read_stream, write_stream) + async with session: + await asyncio.wait_for(session.initialize(), timeout=_get_init_timeout()) + + response = await asyncio.wait_for( + session.list_tools(), timeout=_get_tool_timeout() + ) + + tools = [] + for tool in response.tools: + tools.append( + { + "name": tool.name, + "description": tool.description or "", + "inputSchema": tool.inputSchema if tool.inputSchema else {}, + } + ) + return tools + + +async def _http_call_tool( + base_url: str, + project: ProjectConfig, + tool_name: str, + tool_input: dict[str, Any], + branch_id: str | None = None, +) -> dict[str, Any]: + """Call a tool via HTTP transport (persistent server). + + Args: + base_url: Base URL of the persistent MCP server. + project: Project config for authentication headers. + tool_name: Name of the tool to call. + tool_input: Input arguments for the tool. + branch_id: Optional development branch ID. + + Returns: + Dict with tool result content and error status. + """ + headers = _build_http_headers(project, branch_id) + url = f"{base_url}/mcp" + + async with streamablehttp_client(url=url, headers=headers) as ( + read_stream, + write_stream, + _, + ): + session = ClientSession(read_stream, write_stream) + async with session: + await asyncio.wait_for(session.initialize(), timeout=_get_init_timeout()) + + result = await asyncio.wait_for( + session.call_tool(tool_name, tool_input), + timeout=_get_tool_timeout(), + ) + + return { + "content": _parse_content(result), + "isError": bool(result.isError), + } + + +async def _http_validate_and_call( + base_url: str, + project: ProjectConfig, + tool_name: str, + tool_input: dict[str, Any], + branch_id: str | None = None, +) -> dict[str, Any]: + """Validate and call a tool in one HTTP session (persistent server). + + Args: + base_url: Base URL of the persistent MCP server. + project: Project config for authentication headers. + tool_name: Name of the MCP tool to call. + tool_input: Input arguments for the tool. + branch_id: Optional development branch ID. + + Returns: + Dict with "content" and "isError" keys. + + Raises: + ConfigError: If tool_name not found or required params missing. + """ + headers = _build_http_headers(project, branch_id) + url = f"{base_url}/mcp" + + async with streamablehttp_client(url=url, headers=headers) as ( + read_stream, + write_stream, + _, + ): + session = ClientSession(read_stream, write_stream) + async with session: + await asyncio.wait_for(session.initialize(), timeout=_get_init_timeout()) + + # Step 1: list_tools for validation + response = await asyncio.wait_for( + session.list_tools(), timeout=_get_tool_timeout() + ) + + known_tools = {t.name for t in response.tools} + if tool_name not in known_tools: + raise ConfigError( + f"Unknown MCP tool '{tool_name}'. " + f"Use 'kbagent tool list' to see available tools." + ) + + # Step 2: validate required params + schema: dict[str, Any] = {} + for tool in response.tools: + if tool.name == tool_name: + schema = tool.inputSchema if tool.inputSchema else {} + break + + required = schema.get("required", []) + missing = [param for param in required if param not in tool_input] + + expand_config = AUTO_EXPAND_TOOLS.get(tool_name) + if expand_config: + auto_param = expand_config["param"] + missing = [p for p in missing if p != auto_param] + + if missing: + import json as _json + + params_str = ", ".join(missing) + example_json = _json.dumps({p: "..." for p in missing}) + raise ConfigError( + f"Missing required parameter(s) for '{tool_name}': {params_str}. " + f"Use: kbagent tool call {tool_name} --input '{example_json}'" + ) + + # Step 3: call the tool + result = await asyncio.wait_for( + session.call_tool(tool_name, tool_input), + timeout=_get_tool_timeout(), + ) + + return { + "content": _parse_content(result), + "isError": bool(result.isError), + } + + +async def _http_auto_expand( + base_url: str, + project: ProjectConfig, + tool_name: str, + tool_input: dict[str, Any], + expand_config: dict[str, str], + branch_id: str | None = None, +) -> dict[str, Any]: + """Auto-expand a tool call via HTTP transport (persistent server). + + Same logic as _connect_and_auto_expand but over HTTP. + """ + resolve_tool = expand_config["resolve_tool"] + resolve_key = expand_config["resolve_key"] + param_name = expand_config["param"] + + headers = _build_http_headers(project, branch_id) + url = f"{base_url}/mcp" + + async with streamablehttp_client(url=url, headers=headers) as ( + read_stream, + write_stream, + _, + ): + session = ClientSession(read_stream, write_stream) + async with session: + await asyncio.wait_for(session.initialize(), timeout=_get_init_timeout()) + + # Step 1: Call resolve tool + resolve_result = await asyncio.wait_for( + session.call_tool(resolve_tool, {}), + timeout=_get_tool_timeout(), + ) + + if resolve_result.isError: + return { + "content": _parse_content(resolve_result), + "isError": True, + } + + resolve_items = _parse_content(resolve_result) + item_ids = _extract_ids(resolve_items, resolve_key) + + if not item_ids: + return {"content": [], "isError": False} + + # Step 2: Call target tool for each resolved ID + all_content: list[Any] = [] + has_error = False + + for item_id in item_ids: + call_input = {**tool_input, param_name: item_id} + result = await asyncio.wait_for( + session.call_tool(tool_name, call_input), + timeout=_get_tool_timeout(), + ) + + content = _parse_content(result) + if result.isError: + has_error = True + all_content.extend(content) + + return {"content": all_content, "isError": has_error} + + async def _connect_and_auto_expand( project: ProjectConfig, tool_name: str, @@ -356,13 +710,36 @@ def _extract_ids(content_items: list[Any], key: str) -> list[str]: class McpService(BaseService): """Business logic for MCP tool operations across projects. - Wraps keboola-mcp-server as subprocess via MCP SDK. + Supports two transport modes: + - HTTP (default): Uses a persistent server via McpServerManager. + One server serves all projects with per-request credential headers. + - stdio: Spawns a subprocess per MCP session. Fallback mode. + Read tools execute across all projects in parallel. Write tools target a single project. Uses the same DI pattern as JobService/ConfigService. """ + def _get_server_url(self) -> str | None: + """Get the persistent server URL if HTTP transport is configured. + + Returns: + Base URL string if HTTP transport is active and server is running, + None if stdio mode or server cannot be started. + """ + if _get_transport_mode() != "http": + return None + + try: + from .mcp_transport import get_server_manager + + manager = get_server_manager() + return manager.ensure_running() + except Exception as exc: + logger.warning("Failed to start persistent MCP server, falling back to stdio: %s", exc) + return None + def resolve_project(self, alias: str | None = None) -> tuple[str, ProjectConfig]: """Resolve a single project alias (or the default project). @@ -415,13 +792,20 @@ def list_tools( if not projects: raise ConfigError("No projects configured. Use 'kbagent project add' first.") - # Try each project until one succeeds + # Try HTTP transport first, fall back to stdio + server_url = self._get_server_url() + errors: list[dict[str, str]] = [] for alias, project in projects.items(): try: - tools = asyncio.run( - _connect_and_list_tools(project, branch_id=branch_id) - ) + if server_url: + tools = asyncio.run( + _http_list_tools(server_url, project, branch_id=branch_id) + ) + else: + tools = asyncio.run( + _connect_and_list_tools(project, branch_id=branch_id) + ) # Annotate tools with multi_project flag annotated_tools = [] for tool in tools: @@ -574,6 +958,146 @@ def call_tool( else: return self._call_read_tool(tool_name, tool_input, alias, branch_id=branch_id) + def validate_and_call_tool( + self, + tool_name: str, + tool_input: dict[str, Any] | None = None, + alias: str | None = None, + branch_id: str | None = None, + ) -> dict[str, Any]: + """Validate and call an MCP tool in a single session (no double spawn). + + Opens ONE MCP session: validates tool name + required params from + list_tools(), then calls the tool. This eliminates the separate + validate_tool_input() + call_tool() round-trip. + + Prefers HTTP transport (persistent server) with fallback to stdio. + + For write tools and branch-scoped calls: single project. + For read tools: still runs across all projects in parallel, but each + project's session validates + calls in one go. + + Args: + tool_name: Name of the MCP tool to call. + tool_input: Input arguments for the tool. + alias: Project alias for write tools / single-project mode. + branch_id: Optional development branch ID. Forces single-project mode. + + Returns: + Dict with "results" list and "errors" list. + + Raises: + ConfigError: If tool_name not found or required params missing. + """ + if tool_input is None: + tool_input = {} + + server_url = self._get_server_url() + is_write = _is_write_tool(tool_name) + + # Single-project mode: write tools, branch-scoped, or explicit alias + if branch_id is not None or is_write: + resolved_alias, project = self.resolve_project(alias) + try: + # Check if auto-expand needed + expand_config = AUTO_EXPAND_TOOLS.get(tool_name) + if expand_config and expand_config["param"] not in tool_input: + if server_url: + result = asyncio.run( + _http_auto_expand( + server_url, project, tool_name, tool_input, + expand_config, branch_id=branch_id, + ) + ) + else: + result = asyncio.run( + _connect_and_auto_expand( + project, tool_name, tool_input, expand_config, + branch_id=branch_id, + ) + ) + else: + if server_url: + result = asyncio.run( + _http_validate_and_call( + server_url, project, tool_name, tool_input, + branch_id=branch_id, + ) + ) + else: + result = asyncio.run( + _connect_validate_and_call( + project, tool_name, tool_input, branch_id=branch_id, + ) + ) + result["project_alias"] = resolved_alias + return {"results": [result], "errors": []} + except ConfigError: + raise + except Exception as exc: + return { + "results": [], + "errors": [ + { + "project_alias": resolved_alias, + "error_code": "MCP_ERROR", + "message": str(exc), + } + ], + } + + # Multi-project read: parallel validate+call across all projects + projects = self.resolve_projects([alias]) if alias else self.resolve_projects() + if not projects: + raise ConfigError("No projects configured. Use 'kbagent project add' first.") + + # Check if auto-expand is needed + expand_config = AUTO_EXPAND_TOOLS.get(tool_name) + if expand_config and expand_config["param"] not in tool_input: + return asyncio.run( + self._gather_auto_expand_results( + projects, tool_name, tool_input, expand_config, + branch_id=branch_id, server_url=server_url, + ) + ) + + return asyncio.run( + self._gather_validate_and_call_results( + projects, tool_name, tool_input, + branch_id=branch_id, server_url=server_url, + ) + ) + + async def _gather_validate_and_call_results( + self, + projects: dict[str, ProjectConfig], + tool_name: str, + tool_input: dict[str, Any], + branch_id: str | None = None, + server_url: str | None = None, + ) -> dict[str, Any]: + """Run validate+call across multiple projects in parallel. + + Each project opens one session that validates and calls the tool. + Uses HTTP transport when server_url is available. + """ + max_sessions = _get_max_sessions() + sem = asyncio.Semaphore(max_sessions) if max_sessions > 0 else None + tasks = {} + for a, project in projects.items(): + if server_url: + coro = _http_validate_and_call( + server_url, project, tool_name, tool_input, branch_id=branch_id, + ) + else: + coro = _connect_validate_and_call( + project, tool_name, tool_input, branch_id=branch_id, + ) + if sem is not None: + coro = _semaphored(sem, coro) + tasks[a] = asyncio.create_task(coro) + return await self._gather_results(tasks) + def _call_write_tool( self, tool_name: str, @@ -583,11 +1107,21 @@ def _call_write_tool( ) -> dict[str, Any]: """Execute a write tool on a single project.""" resolved_alias, project = self.resolve_project(alias) + server_url = self._get_server_url() try: - result = asyncio.run( - _connect_and_call_tool(project, tool_name, tool_input, branch_id=branch_id) - ) + if server_url: + result = asyncio.run( + _http_call_tool( + server_url, project, tool_name, tool_input, branch_id=branch_id, + ) + ) + else: + result = asyncio.run( + _connect_and_call_tool( + project, tool_name, tool_input, branch_id=branch_id, + ) + ) result["project_alias"] = resolved_alias return {"results": [result], "errors": []} except Exception as exc: @@ -619,17 +1153,23 @@ def _call_read_tool( if not projects: raise ConfigError("No projects configured. Use 'kbagent project add' first.") + server_url = self._get_server_url() + # Check if auto-expand is needed expand_config = AUTO_EXPAND_TOOLS.get(tool_name) if expand_config and expand_config["param"] not in tool_input: return asyncio.run( self._gather_auto_expand_results( - projects, tool_name, tool_input, expand_config, branch_id=branch_id + projects, tool_name, tool_input, expand_config, + branch_id=branch_id, server_url=server_url, ) ) return asyncio.run( - self._gather_read_results(projects, tool_name, tool_input, branch_id=branch_id) + self._gather_read_results( + projects, tool_name, tool_input, + branch_id=branch_id, server_url=server_url, + ) ) @staticmethod @@ -675,20 +1215,27 @@ async def _gather_auto_expand_results( tool_input: dict[str, Any], expand_config: dict[str, str], branch_id: str | None = None, + server_url: str | None = None, ) -> dict[str, Any]: """Run an auto-expanded tool across multiple projects in parallel. For each project, opens one MCP session, resolves the missing param by calling the resolve tool, then calls the target tool per item. - When KBAGENT_MCP_MAX_SESSIONS is set (> 0), concurrency is throttled. + Uses HTTP transport when server_url is available. """ max_sessions = _get_max_sessions() sem = asyncio.Semaphore(max_sessions) if max_sessions > 0 else None tasks = {} for a, project in projects.items(): - coro = _connect_and_auto_expand( - project, tool_name, tool_input, expand_config, branch_id=branch_id - ) + if server_url: + coro = _http_auto_expand( + server_url, project, tool_name, tool_input, expand_config, + branch_id=branch_id, + ) + else: + coro = _connect_and_auto_expand( + project, tool_name, tool_input, expand_config, branch_id=branch_id, + ) if sem is not None: coro = _semaphored(sem, coro) tasks[a] = asyncio.create_task(coro) @@ -700,16 +1247,24 @@ async def _gather_read_results( tool_name: str, tool_input: dict[str, Any], branch_id: str | None = None, + server_url: str | None = None, ) -> dict[str, Any]: """Run a read tool across multiple projects in parallel using asyncio.gather. - When KBAGENT_MCP_MAX_SESSIONS is set (> 0), concurrency is throttled. + Uses HTTP transport when server_url is available. """ max_sessions = _get_max_sessions() sem = asyncio.Semaphore(max_sessions) if max_sessions > 0 else None tasks = {} for a, project in projects.items(): - coro = _connect_and_call_tool(project, tool_name, tool_input, branch_id=branch_id) + if server_url: + coro = _http_call_tool( + server_url, project, tool_name, tool_input, branch_id=branch_id, + ) + else: + coro = _connect_and_call_tool( + project, tool_name, tool_input, branch_id=branch_id, + ) if sem is not None: coro = _semaphored(sem, coro) tasks[a] = asyncio.create_task(coro) @@ -719,7 +1274,7 @@ def check_server_available(self) -> dict[str, Any]: """Check if MCP server is available (for doctor command). Returns: - Dict with check status and message. + Dict with check status, message, and transport info. """ command = detect_mcp_server_command() if command is None: @@ -735,9 +1290,25 @@ def check_server_available(self) -> dict[str, Any]: ), } + transport_mode = _get_transport_mode() + transport_info = f"transport={transport_mode}" + + # If HTTP mode, check persistent server status + if transport_mode == "http": + try: + from .mcp_transport import get_server_manager + + manager = get_server_manager() + if manager.is_running: + transport_info += f", persistent server running on port {manager.port}" + else: + transport_info += ", persistent server not yet started (lazy start)" + except Exception: + transport_info += ", persistent server unavailable (will fallback to stdio)" + return { "check": "mcp_server", "name": "MCP server", "status": "pass", - "message": f"MCP server available via: {' '.join(command)}", + "message": f"MCP server available via: {' '.join(command)} ({transport_info})", } diff --git a/src/keboola_agent_cli/services/mcp_transport.py b/src/keboola_agent_cli/services/mcp_transport.py new file mode 100644 index 00000000..053c90ad --- /dev/null +++ b/src/keboola_agent_cli/services/mcp_transport.py @@ -0,0 +1,215 @@ +"""Persistent MCP HTTP server manager. + +Manages a single keboola-mcp-server process running in streamable-http mode. +The server stays running between tool calls, eliminating subprocess spawn overhead. + +Per-request project credentials are passed via HTTP headers: +- X-Storage-Token: storage API token +- X-Storage-API-URL: stack URL +- X-Branch-ID: optional branch ID + +One server serves ALL projects - just different headers per request. +""" + +import atexit +import contextlib +import logging +import socket +import subprocess +import time +from typing import Any + +from ..constants import ( + MCP_SERVER_HEALTH_TIMEOUT, + MCP_SERVER_STARTUP_TIMEOUT, +) +from .mcp_service import detect_mcp_server_command + +logger = logging.getLogger(__name__) + + +def _find_free_port() -> int: + """Find a free TCP port by binding to port 0 and reading the assigned port.""" + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +class McpServerManager: + """Manages a persistent keboola-mcp-server process with HTTP transport. + + Singleton-like usage: one instance per CLI process. The server is started + lazily on first use and cleaned up on process exit. + """ + + def __init__(self) -> None: + self._process: subprocess.Popen[bytes] | None = None + self._port: int | None = None + self._base_url: str | None = None + self._registered_atexit: bool = False + + @property + def port(self) -> int | None: + """Return the port the server is listening on, or None if not running.""" + return self._port + + @property + def base_url(self) -> str | None: + """Return the base URL of the running server, or None if not running.""" + return self._base_url + + @property + def is_running(self) -> bool: + """Check if the server process is alive.""" + if self._process is None: + return False + return self._process.poll() is None + + def ensure_running(self) -> str: + """Start the server if not running and return the base URL. + + Returns: + Base URL string like "http://127.0.0.1:PORT". + + Raises: + RuntimeError: If the server cannot be started or fails health check. + """ + if self.is_running and self._base_url is not None: + # Quick health check - if server crashed, restart + if self._health_check(): + return self._base_url + logger.warning("MCP server health check failed, restarting") + self.stop() + + return self._start() + + def _start(self) -> str: + """Start the MCP server process. + + Returns: + Base URL of the running server. + + Raises: + RuntimeError: If the server command is not found or server fails to start. + """ + command_parts = detect_mcp_server_command() + if command_parts is None: + raise RuntimeError( + "Cannot find keboola-mcp-server. " + "Install it with: pip install keboola-mcp-server (or: uvx keboola_mcp_server)" + ) + + port = _find_free_port() + cmd = [ + *command_parts, + "--transport", "streamable-http", + "--host", "127.0.0.1", + "--port", str(port), + ] + + logger.info("Starting persistent MCP server: %s", " ".join(cmd)) + + self._process = subprocess.Popen( + cmd, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + self._port = port + self._base_url = f"http://127.0.0.1:{port}" + + # Register cleanup on exit (only once) + if not self._registered_atexit: + atexit.register(self.stop) + self._registered_atexit = True + + # Wait for server to be ready + if not self._wait_for_ready(): + # Collect stderr for diagnostics + stderr_output = "" + if self._process.stderr: + with contextlib.suppress(Exception): + stderr_output = self._process.stderr.read1(4096).decode(errors="replace") # type: ignore[attr-defined] + self.stop() + raise RuntimeError( + f"MCP server failed to start within {MCP_SERVER_STARTUP_TIMEOUT}s. " + f"Command: {' '.join(cmd)}" + + (f"\nStderr: {stderr_output}" if stderr_output else "") + ) + + logger.info("MCP server ready at %s", self._base_url) + return self._base_url + + def _wait_for_ready(self) -> bool: + """Poll the server until it responds to health check or timeout.""" + deadline = time.monotonic() + MCP_SERVER_STARTUP_TIMEOUT + interval = 0.2 + + while time.monotonic() < deadline: + # Check if process died + if self._process is not None and self._process.poll() is not None: + return False + + if self._health_check(): + return True + + time.sleep(interval) + # Exponential backoff up to 1s + interval = min(interval * 1.5, 1.0) + + return False + + def _health_check(self) -> bool: + """Check if the server is responding via a TCP connection test. + + We use a simple TCP connect instead of HTTP request to avoid + needing httpx as a dependency at this layer. + """ + if self._port is None: + return False + try: + with socket.create_connection( + ("127.0.0.1", self._port), + timeout=MCP_SERVER_HEALTH_TIMEOUT, + ): + return True + except (OSError, ConnectionRefusedError): + return False + + def stop(self) -> None: + """Stop the server process if running.""" + if self._process is not None: + logger.info("Stopping persistent MCP server (pid=%s)", self._process.pid) + try: + self._process.terminate() + try: + self._process.wait(timeout=5) + except subprocess.TimeoutExpired: + self._process.kill() + self._process.wait(timeout=2) + except Exception as exc: + logger.warning("Error stopping MCP server: %s", exc) + finally: + self._process = None + self._port = None + self._base_url = None + + def get_status(self) -> dict[str, Any]: + """Return status info for the doctor command.""" + return { + "running": self.is_running, + "port": self._port, + "base_url": self._base_url, + "pid": self._process.pid if self._process else None, + } + + +# Module-level singleton +_server_manager: McpServerManager | None = None + + +def get_server_manager() -> McpServerManager: + """Get or create the module-level McpServerManager singleton.""" + global _server_manager + if _server_manager is None: + _server_manager = McpServerManager() + return _server_manager diff --git a/tests/conftest.py b/tests/conftest.py index b1382a4c..01d55b7e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -8,6 +8,12 @@ from keboola_agent_cli.output import OutputFormatter +@pytest.fixture(autouse=True) +def _force_stdio_transport(monkeypatch: pytest.MonkeyPatch) -> None: + """Force stdio transport in all tests to prevent spawning persistent server.""" + monkeypatch.setenv("KBAGENT_MCP_TRANSPORT", "stdio") + + @pytest.fixture def tmp_config_dir(tmp_path: Path) -> Path: """Provide a temporary directory for configuration files.""" diff --git a/tests/test_cli.py b/tests/test_cli.py index 5c72e762..3c2ff3fe 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -3220,8 +3220,7 @@ def test_tool_call_read_json(self, tmp_path: Path) -> None: MockJobService.return_value = JobService(config_store=store) mock_mcp = MagicMock() - mock_mcp.validate_tool_input.return_value = ([], {"list_configs", "get_config", "create_config"}) - mock_mcp.call_tool.return_value = SAMPLE_TOOL_RESULT_MULTI + mock_mcp.validate_and_call_tool.return_value = SAMPLE_TOOL_RESULT_MULTI MockMcpService.return_value = mock_mcp result = runner.invoke( @@ -3237,12 +3236,11 @@ def test_tool_call_read_json(self, tmp_path: Path) -> None: assert results[0]["project_alias"] == "prod" assert results[1]["project_alias"] == "dev" assert results[0]["isError"] is False - mock_mcp.call_tool.assert_called_once_with( + mock_mcp.validate_and_call_tool.assert_called_once_with( tool_name="list_configs", tool_input={}, alias=None, branch_id=None, - _known_tools={"list_configs", "get_config", "create_config"}, ) def test_tool_call_write_json(self, tmp_path: Path) -> None: @@ -3270,8 +3268,7 @@ def test_tool_call_write_json(self, tmp_path: Path) -> None: MockJobService.return_value = JobService(config_store=store) mock_mcp = MagicMock() - mock_mcp.validate_tool_input.return_value = ([], {"list_configs", "get_config", "create_config"}) - mock_mcp.call_tool.return_value = SAMPLE_TOOL_RESULT + mock_mcp.validate_and_call_tool.return_value = SAMPLE_TOOL_RESULT MockMcpService.return_value = mock_mcp result = runner.invoke( @@ -3295,12 +3292,11 @@ def test_tool_call_write_json(self, tmp_path: Path) -> None: assert len(results) == 1 assert results[0]["project_alias"] == "prod" assert results[0]["isError"] is False - mock_mcp.call_tool.assert_called_once_with( + mock_mcp.validate_and_call_tool.assert_called_once_with( tool_name="create_config", tool_input={"name": "New Config", "component_id": "keboola.ex-db-snowflake"}, alias="prod", branch_id=None, - _known_tools={"list_configs", "get_config", "create_config"}, ) def test_tool_call_invalid_input(self, tmp_path: Path) -> None: @@ -3347,7 +3343,7 @@ def test_tool_call_invalid_input(self, tmp_path: Path) -> None: assert output["status"] == "error" assert output["error"]["code"] == "INVALID_ARGUMENT" assert "Invalid JSON" in output["error"]["message"] - mock_mcp.call_tool.assert_not_called() + mock_mcp.validate_and_call_tool.assert_not_called() def test_tool_call_config_error(self, tmp_path: Path) -> None: """tool call when no projects configured returns exit code 5.""" @@ -3369,7 +3365,7 @@ def test_tool_call_config_error(self, tmp_path: Path) -> None: MockJobService.return_value = JobService(config_store=store) mock_mcp = MagicMock() - mock_mcp.validate_tool_input.side_effect = ConfigError( + mock_mcp.validate_and_call_tool.side_effect = ConfigError( "No projects configured. Use 'kbagent project add' first." ) MockMcpService.return_value = mock_mcp @@ -3410,8 +3406,7 @@ def test_tool_call_human_output(self, tmp_path: Path) -> None: MockJobService.return_value = JobService(config_store=store) mock_mcp = MagicMock() - mock_mcp.validate_tool_input.return_value = ([], {"list_configs", "get_config", "create_config"}) - mock_mcp.call_tool.return_value = SAMPLE_TOOL_RESULT + mock_mcp.validate_and_call_tool.return_value = SAMPLE_TOOL_RESULT MockMcpService.return_value = mock_mcp result = runner.invoke( @@ -4698,8 +4693,7 @@ def test_tool_call_branch_with_project_ok(self, tmp_path: Path) -> None: MockJobService.return_value = JobService(config_store=store) mock_mcp = MagicMock() - mock_mcp.validate_tool_input.return_value = ([], {"list_configs", "get_config", "create_config"}) - mock_mcp.call_tool.return_value = { + mock_mcp.validate_and_call_tool.return_value = { "results": [ { "content": [{"configs": ["cfg1"]}], @@ -4720,12 +4714,11 @@ def test_tool_call_branch_with_project_ok(self, tmp_path: Path) -> None: output = json.loads(result.output) assert output["status"] == "ok" # Verify branch_id was passed to the service - mock_mcp.call_tool.assert_called_once_with( + mock_mcp.validate_and_call_tool.assert_called_once_with( tool_name="list_configs", tool_input={}, alias="prod", branch_id="456", - _known_tools={"list_configs", "get_config", "create_config"}, ) @@ -5252,8 +5245,7 @@ def test_tool_call_auto_resolves_active_branch(self, tmp_path: Path) -> None: MockJobService.return_value = JobService(config_store=store) mock_mcp = MagicMock() - mock_mcp.validate_tool_input.return_value = ([], {"list_configs", "get_config", "create_config"}) - mock_mcp.call_tool.return_value = { + mock_mcp.validate_and_call_tool.return_value = { "results": [ { "content": [{"configs": ["cfg1"]}], @@ -5275,10 +5267,9 @@ def test_tool_call_auto_resolves_active_branch(self, tmp_path: Path) -> None: output = json.loads(result.output) assert output["status"] == "ok" # Verify the service was called with auto-resolved branch - mock_mcp.call_tool.assert_called_once_with( + mock_mcp.validate_and_call_tool.assert_called_once_with( tool_name="list_configs", tool_input={}, alias="prod", branch_id="456", - _known_tools={"list_configs", "get_config", "create_config"}, ) diff --git a/tests/test_mcp_service.py b/tests/test_mcp_service.py index 102af73d..de6508f1 100644 --- a/tests/test_mcp_service.py +++ b/tests/test_mcp_service.py @@ -144,35 +144,49 @@ class TestDetectMcpServerCommand: """Tests for detect_mcp_server_command() which finds the MCP server binary.""" @patch("keboola_agent_cli.services.mcp_service.shutil.which") - def test_uvx_available(self, mock_which: MagicMock) -> None: - """When uvx is available, returns ['uvx', '--prerelease=allow', 'keboola_mcp_server@latest'].""" - mock_which.side_effect = lambda cmd: "/usr/local/bin/uvx" if cmd == "uvx" else None + def test_local_install_preferred(self, mock_which: MagicMock) -> None: + """When keboola_mcp_server is locally installed, returns it (fastest).""" + mock_which.side_effect = lambda cmd: "/usr/local/bin/keboola_mcp_server" if cmd == "keboola_mcp_server" else None result = detect_mcp_server_command() - assert result == ["uvx", "--prerelease=allow", "keboola_mcp_server@latest"] + assert result == ["keboola_mcp_server"] + @patch("keboola_agent_cli.services.mcp_service.subprocess.run") @patch("keboola_agent_cli.services.mcp_service.shutil.which") - def test_keboola_mcp_server_available(self, mock_which: MagicMock) -> None: - """When uvx is not available but keboola_mcp_server is, returns it directly.""" + def test_python_module_second(self, mock_which: MagicMock, mock_run: MagicMock) -> None: + """When python is available and module exists, returns python -m.""" def which_side_effect(cmd: str) -> str | None: - if cmd == "keboola_mcp_server": - return "/usr/local/bin/keboola_mcp_server" + if cmd == "python": + return "/usr/bin/python" return None mock_which.side_effect = which_side_effect + mock_run.return_value = MagicMock(returncode=0) result = detect_mcp_server_command() - assert result == ["keboola_mcp_server"] + assert result == ["python", "-m", "keboola_mcp_server"] + mock_run.assert_called_once() + @patch("keboola_agent_cli.services.mcp_service.subprocess.run") @patch("keboola_agent_cli.services.mcp_service.shutil.which") - def test_python_fallback(self, mock_which: MagicMock) -> None: - """When only python is available, returns python -m fallback.""" + def test_python_module_not_installed_falls_to_uvx( + self, mock_which: MagicMock, mock_run: MagicMock, + ) -> None: + """When python exists but module not installed, falls back to uvx.""" def which_side_effect(cmd: str) -> str | None: - if cmd == "python": - return "/usr/bin/python" + if cmd in ("python", "uvx"): + return f"/usr/bin/{cmd}" return None mock_which.side_effect = which_side_effect + mock_run.return_value = MagicMock(returncode=1) result = detect_mcp_server_command() - assert result == ["python", "-m", "keboola_mcp_server"] + assert result == ["uvx", "--prerelease=allow", "keboola_mcp_server"] + + @patch("keboola_agent_cli.services.mcp_service.shutil.which") + def test_uvx_fallback(self, mock_which: MagicMock) -> None: + """When only uvx is available, returns uvx WITHOUT @latest.""" + mock_which.side_effect = lambda cmd: "/usr/local/bin/uvx" if cmd == "uvx" else None + result = detect_mcp_server_command() + assert result == ["uvx", "--prerelease=allow", "keboola_mcp_server"] @patch("keboola_agent_cli.services.mcp_service.shutil.which") def test_nothing_available(self, mock_which: MagicMock) -> None: @@ -181,6 +195,18 @@ def test_nothing_available(self, mock_which: MagicMock) -> None: result = detect_mcp_server_command() assert result is None + @patch("keboola_agent_cli.services.mcp_service.shutil.which") + def test_local_install_preferred_over_uvx(self, mock_which: MagicMock) -> None: + """Local install is preferred even when uvx is also available.""" + def which_side_effect(cmd: str) -> str | None: + if cmd in ("keboola_mcp_server", "uvx"): + return f"/usr/local/bin/{cmd}" + return None + + mock_which.side_effect = which_side_effect + result = detect_mcp_server_command() + assert result == ["keboola_mcp_server"] + # --------------------------------------------------------------------------- # TestMcpServiceResolveProject @@ -525,10 +551,14 @@ def test_call_write_tool_no_default_raises(self, tmp_path: Path) -> None: svc.call_tool("create_config", {"name": "new-config"}) @patch("keboola_agent_cli.services.mcp_service._connect_and_call_tool", new_callable=AsyncMock) + @patch("keboola_agent_cli.services.mcp_service._connect_and_list_tools", new_callable=AsyncMock) def test_call_tool_none_input_defaults_to_empty_dict( - self, mock_call_tool: AsyncMock, tmp_path: Path + self, mock_list_tools: AsyncMock, mock_call_tool: AsyncMock, tmp_path: Path ) -> None: """When tool_input is None, it defaults to an empty dict.""" + mock_list_tools.return_value = [ + {"name": "list_configs", "description": "List configs", "inputSchema": {}}, + ] mock_call_tool.return_value = { "content": [{"result": "ok"}], "isError": False, @@ -556,15 +586,16 @@ class TestMcpServiceErrorAccumulation: """Tests for error accumulation - projects that fail don't stop others.""" @patch("keboola_agent_cli.services.mcp_service._connect_and_call_tool", new_callable=AsyncMock) + @patch("keboola_agent_cli.services.mcp_service._connect_and_list_tools", new_callable=AsyncMock) def test_read_tool_partial_failure( - self, mock_call_tool: AsyncMock, tmp_path: Path + self, mock_list_tools: AsyncMock, mock_call_tool: AsyncMock, tmp_path: Path ) -> None: """When some projects fail for a read tool, successful results and errors are both returned.""" - call_count = 0 + mock_list_tools.return_value = [ + {"name": "list_configs", "description": "List configs", "inputSchema": {}}, + ] async def side_effect(project, tool_name, tool_input, branch_id=None): - nonlocal call_count - call_count += 1 if project.token == "tok-failing": raise RuntimeError("Connection timeout") return { @@ -657,17 +688,19 @@ def which_side_effect(cmd: str) -> str | None: assert result["status"] == "pass" assert "keboola_mcp_server" in result["message"] + @patch("keboola_agent_cli.services.mcp_service.subprocess.run") @patch("keboola_agent_cli.services.mcp_service.shutil.which") def test_server_available_via_python( - self, mock_which: MagicMock, tmp_path: Path + self, mock_which: MagicMock, mock_run: MagicMock, tmp_path: Path ) -> None: - """When only python is available, status is 'pass' with python -m command.""" + """When only python is available and module exists, status is 'pass'.""" def which_side_effect(cmd: str) -> str | None: if cmd == "python": return "/usr/bin/python" return None mock_which.side_effect = which_side_effect + mock_run.return_value = MagicMock(returncode=0) store = _setup_store(tmp_path, projects={}) svc = McpService(config_store=store) diff --git a/tests/test_mcp_transport.py b/tests/test_mcp_transport.py new file mode 100644 index 00000000..eb9289b9 --- /dev/null +++ b/tests/test_mcp_transport.py @@ -0,0 +1,265 @@ +"""Tests for McpServerManager - persistent MCP HTTP server manager.""" + +import socket +from unittest.mock import MagicMock, patch + +import pytest + +from keboola_agent_cli.services.mcp_transport import ( + McpServerManager, + _find_free_port, + get_server_manager, +) + + +class TestFindFreePort: + """Tests for _find_free_port().""" + + def test_returns_valid_port(self) -> None: + """Should return a port number in valid range.""" + port = _find_free_port() + assert isinstance(port, int) + assert 1024 <= port <= 65535 + + def test_returns_different_ports(self) -> None: + """Consecutive calls should return different ports (usually).""" + ports = {_find_free_port() for _ in range(5)} + # At least 2 different ports out of 5 calls + assert len(ports) >= 2 + + +class TestMcpServerManager: + """Tests for McpServerManager lifecycle.""" + + def test_initial_state(self) -> None: + """New manager starts with no server running.""" + manager = McpServerManager() + assert manager.is_running is False + assert manager.port is None + assert manager.base_url is None + + def test_get_status_not_running(self) -> None: + """Status when server is not running.""" + manager = McpServerManager() + status = manager.get_status() + assert status["running"] is False + assert status["port"] is None + assert status["base_url"] is None + assert status["pid"] is None + + @patch("keboola_agent_cli.services.mcp_transport.detect_mcp_server_command") + def test_start_no_command_raises(self, mock_detect: MagicMock) -> None: + """If no MCP server command is found, RuntimeError is raised.""" + mock_detect.return_value = None + manager = McpServerManager() + + with pytest.raises(RuntimeError, match="Cannot find keboola-mcp-server"): + manager.ensure_running() + + @patch("keboola_agent_cli.services.mcp_transport.detect_mcp_server_command") + def test_start_server_fails_to_respond(self, mock_detect: MagicMock) -> None: + """If server process starts but never becomes healthy, RuntimeError is raised.""" + mock_detect.return_value = ["echo", "test"] + + manager = McpServerManager() + + # Mock _wait_for_ready to always return False (timeout) + with ( + patch.object(manager, "_wait_for_ready", return_value=False), + patch("subprocess.Popen") as mock_popen, + ): + mock_process = MagicMock() + mock_process.poll.return_value = None + mock_process.stderr = MagicMock() + mock_process.stderr.read1.return_value = b"some error" + mock_process.pid = 12345 + mock_popen.return_value = mock_process + + with pytest.raises(RuntimeError, match="MCP server failed to start"): + manager.ensure_running() + + @patch("keboola_agent_cli.services.mcp_transport.detect_mcp_server_command") + def test_start_and_stop(self, mock_detect: MagicMock) -> None: + """Starting and stopping the server cleans up state.""" + mock_detect.return_value = ["echo", "test"] + + manager = McpServerManager() + + with ( + patch.object(manager, "_wait_for_ready", return_value=True), + patch("subprocess.Popen") as mock_popen, + ): + mock_process = MagicMock() + mock_process.poll.return_value = None + mock_process.pid = 12345 + mock_popen.return_value = mock_process + + url = manager.ensure_running() + assert url.startswith("http://127.0.0.1:") + assert manager.is_running is True + assert manager.port is not None + + manager.stop() + mock_process.terminate.assert_called_once() + assert manager.is_running is False + assert manager.port is None + assert manager.base_url is None + + @patch("keboola_agent_cli.services.mcp_transport.detect_mcp_server_command") + def test_ensure_running_reuses_existing(self, mock_detect: MagicMock) -> None: + """If server is already running and healthy, reuse it.""" + mock_detect.return_value = ["echo", "test"] + + manager = McpServerManager() + + with ( + patch.object(manager, "_wait_for_ready", return_value=True), + patch("subprocess.Popen") as mock_popen, + ): + mock_process = MagicMock() + mock_process.poll.return_value = None + mock_process.pid = 12345 + mock_popen.return_value = mock_process + + url1 = manager.ensure_running() + + with patch.object(manager, "_health_check", return_value=True): + url2 = manager.ensure_running() + + assert url1 == url2 + # Popen should only be called once (reuse) + assert mock_popen.call_count == 1 + + manager.stop() + + def test_health_check_no_port(self) -> None: + """Health check with no port returns False.""" + manager = McpServerManager() + assert manager._health_check() is False + + def test_health_check_connection_refused(self) -> None: + """Health check to closed port returns False.""" + manager = McpServerManager() + manager._port = _find_free_port() + assert manager._health_check() is False + + def test_health_check_real_server(self) -> None: + """Health check to a real listening socket returns True.""" + manager = McpServerManager() + + # Start a temporary TCP server + srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + srv.bind(("127.0.0.1", 0)) + srv.listen(1) + port = srv.getsockname()[1] + + try: + manager._port = port + assert manager._health_check() is True + finally: + srv.close() + + def test_stop_when_not_running(self) -> None: + """Stopping when no server is running is a no-op.""" + manager = McpServerManager() + manager.stop() # Should not raise + + +class TestGetServerManager: + """Tests for the module-level singleton.""" + + def test_returns_same_instance(self) -> None: + """get_server_manager() returns the same instance.""" + m1 = get_server_manager() + m2 = get_server_manager() + assert m1 is m2 + + def test_returns_mcpservermanager_instance(self) -> None: + """get_server_manager() returns an McpServerManager.""" + m = get_server_manager() + assert isinstance(m, McpServerManager) + + +class TestHttpTransportFunctions: + """Tests for HTTP transport helper functions in mcp_service.""" + + def test_get_transport_mode_default(self, monkeypatch: pytest.MonkeyPatch) -> None: + """Default transport mode is 'http'.""" + monkeypatch.delenv("KBAGENT_MCP_TRANSPORT", raising=False) + from keboola_agent_cli.services.mcp_service import _get_transport_mode + + assert _get_transport_mode() == "http" + + def test_get_transport_mode_stdio(self, monkeypatch: pytest.MonkeyPatch) -> None: + """KBAGENT_MCP_TRANSPORT=stdio returns 'stdio'.""" + monkeypatch.setenv("KBAGENT_MCP_TRANSPORT", "stdio") + from keboola_agent_cli.services.mcp_service import _get_transport_mode + + assert _get_transport_mode() == "stdio" + + def test_build_http_headers(self) -> None: + """Headers include token and stack URL.""" + from keboola_agent_cli.models import ProjectConfig + from keboola_agent_cli.services.mcp_service import _build_http_headers + + project = ProjectConfig( + stack_url="https://connection.keboola.com", + token="test-token", + ) + headers = _build_http_headers(project) + assert headers["X-Storage-Token"] == "test-token" + assert headers["X-Storage-API-URL"] == "https://connection.keboola.com" + assert "X-Branch-ID" not in headers + + def test_build_http_headers_with_branch(self) -> None: + """Headers include branch ID when provided.""" + from keboola_agent_cli.models import ProjectConfig + from keboola_agent_cli.services.mcp_service import _build_http_headers + + project = ProjectConfig( + stack_url="https://connection.keboola.com", + token="test-token", + ) + headers = _build_http_headers(project, branch_id="123") + assert headers["X-Branch-ID"] == "123" + + +class TestMcpServiceTransportSelection: + """Tests for McpService._get_server_url() transport selection.""" + + def test_stdio_mode_returns_none( + self, tmp_path, monkeypatch: pytest.MonkeyPatch + ) -> None: + """In stdio mode, _get_server_url() returns None.""" + monkeypatch.setenv("KBAGENT_MCP_TRANSPORT", "stdio") + from keboola_agent_cli.config_store import ConfigStore + from keboola_agent_cli.services.mcp_service import McpService + + config_dir = tmp_path / "config" + config_dir.mkdir() + store = ConfigStore(config_dir=config_dir) + svc = McpService(config_store=store) + + assert svc._get_server_url() is None + + def test_http_mode_with_failed_server_returns_none( + self, tmp_path, monkeypatch: pytest.MonkeyPatch + ) -> None: + """In HTTP mode, if server fails to start, returns None (fallback to stdio).""" + monkeypatch.setenv("KBAGENT_MCP_TRANSPORT", "http") + from keboola_agent_cli.config_store import ConfigStore + from keboola_agent_cli.services.mcp_service import McpService + + config_dir = tmp_path / "config" + config_dir.mkdir() + store = ConfigStore(config_dir=config_dir) + svc = McpService(config_store=store) + + mock_manager = MagicMock() + mock_manager.ensure_running.side_effect = RuntimeError("No server") + + with patch( + "keboola_agent_cli.services.mcp_transport.get_server_manager", + return_value=mock_manager, + ): + assert svc._get_server_url() is None From 76af4841845bee41e5f56728eceb325735311f26 Mon Sep 17 00:00:00 2001 From: Petr Date: Thu, 5 Mar 2026 15:18:36 +0100 Subject: [PATCH 2/3] Bump version to 0.6.6 Co-Authored-By: Claude Opus 4.6 --- src/keboola_agent_cli/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/keboola_agent_cli/__init__.py b/src/keboola_agent_cli/__init__.py index 0d7df30f..3e04ff95 100644 --- a/src/keboola_agent_cli/__init__.py +++ b/src/keboola_agent_cli/__init__.py @@ -1,3 +1,3 @@ """Keboola Agent CLI - AI-friendly interface to Keboola projects.""" -__version__ = "0.6.5" +__version__ = "0.6.6" From c9d4e63b4c69b3b0e1636c34adbab3255361d7bd Mon Sep 17 00:00:00 2001 From: Petr Date: Thu, 5 Mar 2026 15:33:07 +0100 Subject: [PATCH 3/3] Fix default transport to stdio for multi-project compatibility HTTP transport fails with 34 concurrent MCP sessions. Default to stdio which works reliably and still benefits from Phase 1 optimizations (single-session validate+call, detection reorder: ~11s vs ~65s). HTTP remains available via KBAGENT_MCP_TRANSPORT=http for single-project. Co-Authored-By: Claude Opus 4.6 --- src/keboola_agent_cli/constants.py | 2 +- tests/test_mcp_transport.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/src/keboola_agent_cli/constants.py b/src/keboola_agent_cli/constants.py index 5e55b5e0..736585e4 100644 --- a/src/keboola_agent_cli/constants.py +++ b/src/keboola_agent_cli/constants.py @@ -44,7 +44,7 @@ # --- MCP HTTP Transport --- # Transport mode: "http" (persistent server) or "stdio" (subprocess per call) ENV_MCP_TRANSPORT: str = "KBAGENT_MCP_TRANSPORT" -DEFAULT_MCP_TRANSPORT: str = "http" +DEFAULT_MCP_TRANSPORT: str = "stdio" # Timeout for the persistent MCP server to start and be healthy MCP_SERVER_STARTUP_TIMEOUT: float = 15.0 # Timeout for health check requests to persistent MCP server diff --git a/tests/test_mcp_transport.py b/tests/test_mcp_transport.py index eb9289b9..2edf18cc 100644 --- a/tests/test_mcp_transport.py +++ b/tests/test_mcp_transport.py @@ -184,11 +184,11 @@ class TestHttpTransportFunctions: """Tests for HTTP transport helper functions in mcp_service.""" def test_get_transport_mode_default(self, monkeypatch: pytest.MonkeyPatch) -> None: - """Default transport mode is 'http'.""" + """Default transport mode is 'stdio'.""" monkeypatch.delenv("KBAGENT_MCP_TRANSPORT", raising=False) from keboola_agent_cli.services.mcp_service import _get_transport_mode - assert _get_transport_mode() == "http" + assert _get_transport_mode() == "stdio" def test_get_transport_mode_stdio(self, monkeypatch: pytest.MonkeyPatch) -> None: """KBAGENT_MCP_TRANSPORT=stdio returns 'stdio'."""