From 2f9707cfdf8fcf631f990843d274ed20ff1eddf5 Mon Sep 17 00:00:00 2001 From: Aditya Chhabra Date: Wed, 25 Feb 2026 14:17:34 +0000 Subject: [PATCH] refactor(mcp): fix mock placement, type annotation, and constructor validation - Extract MockMCPTransport to features/mocks/mock_mcp_transport.py - Use ToolRegistry type with TYPE_CHECKING in register_tools() - Move transport config validation into __init__ via _validate_config() --- features/mocks/__init__.py | 3 +- features/mocks/mock_mcp_transport.py | 68 ++++++++++++++++++++++++++++ features/steps/mcp_adapter_steps.py | 63 +------------------------- robot/helper_mcp_adapter.py | 47 +++++-------------- src/cleveragents/mcp/adapter.py | 30 +++++++----- 5 files changed, 102 insertions(+), 109 deletions(-) create mode 100644 features/mocks/mock_mcp_transport.py diff --git a/features/mocks/__init__.py b/features/mocks/__init__.py index 33b3d3632..038b9843c 100644 --- a/features/mocks/__init__.py +++ b/features/mocks/__init__.py @@ -6,5 +6,6 @@ features/ directory and never in production code (implementation_plan.md). """ from .mock_ai_provider import MockAIProvider +from .mock_mcp_transport import MockMCPTransport -__all__ = ["MockAIProvider"] +__all__ = ["MockAIProvider", "MockMCPTransport"] diff --git a/features/mocks/mock_mcp_transport.py b/features/mocks/mock_mcp_transport.py new file mode 100644 index 000000000..e40409e7a --- /dev/null +++ b/features/mocks/mock_mcp_transport.py @@ -0,0 +1,68 @@ +"""Shared mock MCP transport for Behave and Robot Framework tests.""" + +from __future__ import annotations + +from typing import Any + +from cleveragents.mcp.adapter import MCPServerConfig, MCPTransport + + +class MockMCPTransport(MCPTransport): + """In-memory mock transport simulating an MCP server. + + Supports configurable tools, connection failures, invocation results, + invocation errors, and tool-level timeouts for deterministic testing. + """ + + def __init__( + self, + tools: list[dict[str, Any]] | None = None, + *, + fail_connect: bool = False, + invoke_results: dict[str, dict[str, Any]] | None = None, + invoke_errors: dict[str, str] | None = None, + timeout_tools: set[str] | None = None, + ) -> None: + self._tools = tools or [] + self._fail_connect = fail_connect + self._invoke_results = invoke_results or {} + self._invoke_errors = invoke_errors or {} + self._timeout_tools = timeout_tools or set() + self._connected = False + + def connect(self, config: MCPServerConfig) -> dict[str, Any]: + if self._fail_connect: + msg = "Mock connection refused" + raise ConnectionRefusedError(msg) + self._connected = True + return {"capabilities": {"tools": True}} + + def call(self, method: str, params: dict[str, Any]) -> dict[str, Any]: + if method == "tools/list": + return {"tools": list(self._tools)} + + if method == "tools/call": + tool_name = params.get("name", "") + + if tool_name in self._timeout_tools: + msg = f"Tool '{tool_name}' exceeded timeout" + raise TimeoutError(msg) + + if tool_name in self._invoke_errors: + return { + "isError": True, + "error": self._invoke_errors[tool_name], + } + + if tool_name in self._invoke_results: + return {"content": self._invoke_results[tool_name]} + + return {"content": {"result": "ok"}} + + return {} + + def close(self) -> None: + self._connected = False + + def add_tool(self, tool: dict[str, Any]) -> None: + self._tools.append(tool) diff --git a/features/steps/mcp_adapter_steps.py b/features/steps/mcp_adapter_steps.py index 62af05f8b..6b176e153 100644 --- a/features/steps/mcp_adapter_steps.py +++ b/features/steps/mcp_adapter_steps.py @@ -13,70 +13,9 @@ from cleveragents.mcp.adapter import ( MCPServerConfig, MCPToolAdapter, MCPToolFilter, - MCPTransport, ) from cleveragents.tool.registry import ToolRegistry - -# --------------------------------------------------------------- -# Mock transport for testing -# --------------------------------------------------------------- - - -class MockMCPTransport(MCPTransport): - """In-memory mock transport simulating an MCP server.""" - - def __init__( - self, - tools: list[dict[str, Any]] | None = None, - *, - fail_connect: bool = False, - invoke_results: dict[str, dict[str, Any]] | None = None, - invoke_errors: dict[str, str] | None = None, - timeout_tools: set[str] | None = None, - ) -> None: - self._tools = tools or [] - self._fail_connect = fail_connect - self._invoke_results = invoke_results or {} - self._invoke_errors = invoke_errors or {} - self._timeout_tools = timeout_tools or set() - self._connected = False - - def connect(self, config: MCPServerConfig) -> dict[str, Any]: - if self._fail_connect: - msg = "Mock connection refused" - raise ConnectionRefusedError(msg) - self._connected = True - return {"capabilities": {"tools": True}} - - def call(self, method: str, params: dict[str, Any]) -> dict[str, Any]: - if method == "tools/list": - return {"tools": list(self._tools)} - - if method == "tools/call": - tool_name = params.get("name", "") - - if tool_name in self._timeout_tools: - msg = f"Tool '{tool_name}' exceeded timeout" - raise TimeoutError(msg) - - if tool_name in self._invoke_errors: - return { - "isError": True, - "error": self._invoke_errors[tool_name], - } - - if tool_name in self._invoke_results: - return {"content": self._invoke_results[tool_name]} - - return {"content": {"result": "ok"}} - - return {} - - def close(self) -> None: - self._connected = False - - def add_tool(self, tool: dict[str, Any]) -> None: - self._tools.append(tool) +from features.mocks.mock_mcp_transport import MockMCPTransport def _mock_tool(name: str, desc: str = "", schema: dict | None = None) -> dict: diff --git a/robot/helper_mcp_adapter.py b/robot/helper_mcp_adapter.py index 0be4ded66..c2dd6804f 100644 --- a/robot/helper_mcp_adapter.py +++ b/robot/helper_mcp_adapter.py @@ -14,42 +14,18 @@ import sys from pathlib import Path from typing import Any +_ROOT = str(Path(__file__).resolve().parents[1]) _SRC = str(Path(__file__).resolve().parents[1] / "src") -if _SRC not in sys.path: - sys.path.insert(0, _SRC) +for _p in (_SRC, _ROOT): + if _p not in sys.path: + sys.path.insert(0, _p) -from cleveragents.mcp.adapter import ( # noqa: E402 - MCPServerConfig, - MCPToolAdapter, - MCPTransport, -) +from features.mocks.mock_mcp_transport import MockMCPTransport # noqa: E402 + +from cleveragents.mcp.adapter import MCPServerConfig, MCPToolAdapter # noqa: E402 from cleveragents.tool.registry import ToolRegistry # noqa: E402 -class _MockTransport(MCPTransport): - """Minimal mock transport for smoke tests.""" - - def __init__(self, tools: list[dict[str, Any]] | None = None) -> None: - self._tools = tools or [] - self._results: dict[str, dict[str, Any]] = {} - - def connect(self, config: MCPServerConfig) -> dict[str, Any]: - return {"capabilities": {"tools": True}} - - def call(self, method: str, params: dict[str, Any]) -> dict[str, Any]: - if method == "tools/list": - return {"tools": list(self._tools)} - if method == "tools/call": - name = params.get("name", "") - if name in self._results: - return {"content": self._results[name]} - return {"content": {"result": "ok"}} - return {} - - def close(self) -> None: - pass - - def _mock_tool(name: str) -> dict[str, Any]: return {"name": name, "description": f"Mock {name}", "inputSchema": {}} @@ -89,7 +65,7 @@ def _discover_tools() -> int: _mock_tool("search"), ] config = MCPServerConfig(name="test", transport="stdio", command="echo") - transport = _MockTransport(tools=tools) + transport = MockMCPTransport(tools=tools) adapter = MCPToolAdapter(config=config, transport=transport) adapter.connect() discovered = adapter.discover_tools() @@ -102,8 +78,9 @@ def _discover_tools() -> int: def _invoke_tool() -> int: tools = [_mock_tool("create_issue")] config = MCPServerConfig(name="test", transport="stdio", command="echo") - transport = _MockTransport(tools=tools) - transport._results["create_issue"] = {"id": 42} + transport = MockMCPTransport( + tools=tools, invoke_results={"create_issue": {"id": 42}} + ) adapter = MCPToolAdapter(config=config, transport=transport) adapter.connect() adapter.discover_tools() @@ -116,7 +93,7 @@ def _invoke_tool() -> int: def _register_tools() -> int: tools = [_mock_tool("tool_a"), _mock_tool("tool_b")] config = MCPServerConfig(name="test", transport="stdio", command="echo") - transport = _MockTransport(tools=tools) + transport = MockMCPTransport(tools=tools) adapter = MCPToolAdapter(config=config, transport=transport) adapter.connect() registry = ToolRegistry() diff --git a/src/cleveragents/mcp/adapter.py b/src/cleveragents/mcp/adapter.py index 7d932567a..9057c4ba0 100644 --- a/src/cleveragents/mcp/adapter.py +++ b/src/cleveragents/mcp/adapter.py @@ -18,7 +18,10 @@ from __future__ import annotations import logging import threading import time -from typing import Any +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from cleveragents.tool.registry import ToolRegistry import jsonschema from pydantic import BaseModel, ConfigDict, Field @@ -138,6 +141,7 @@ class MCPToolAdapter: config: MCPServerConfig, transport: MCPTransport | None = None, ) -> None: + self._validate_config(config) self._config = config self._transport = transport or MCPTransport() self._connected = False @@ -145,15 +149,11 @@ class MCPToolAdapter: self._tools: dict[str, MCPToolDescriptor] = {} self._lock = threading.RLock() - @classmethod - def from_server_config( - cls, - config: MCPServerConfig, - transport: MCPTransport | None = None, - ) -> MCPToolAdapter: - """Create an adapter with config validation. + @staticmethod + def _validate_config(config: MCPServerConfig) -> None: + """Validate transport-specific required fields. - Raises ``ValueError`` when required transport fields are missing. + Raises ``ValueError`` when required fields are missing. """ if config.transport == "stdio" and not config.command: msg = ( @@ -167,6 +167,14 @@ class MCPToolAdapter: f"transport requires 'url' field" ) raise ValueError(msg) + + @classmethod + def from_server_config( + cls, + config: MCPServerConfig, + transport: MCPTransport | None = None, + ) -> MCPToolAdapter: + """Create an adapter (validation is performed in ``__init__``).""" return cls(config=config, transport=transport) @property @@ -400,7 +408,7 @@ class MCPToolAdapter: def register_tools( self, - registry: Any, + registry: ToolRegistry, namespace: str, tool_filter: MCPToolFilter | None = None, ) -> list[str]: @@ -409,7 +417,7 @@ class MCPToolAdapter: Parameters ---------- registry: - A ``ToolRegistry`` instance. + The ``ToolRegistry`` to register tools into. namespace: Namespace prefix for registered tool names. tool_filter: