forked from cleveragents/cleveragents-core
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()
This commit is contained in:
@@ -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"]
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
+12
-35
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user