Skip to content

Commit f49692a

Browse files
committed
feature: mcp tool加上缓存避免多次网络访问 (#71)
1 parent cf5e977 commit f49692a

3 files changed

Lines changed: 297 additions & 4 deletions

File tree

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@ classifiers = [
2626
dependencies = [
2727
"pydantic>=2.11.3",
2828
"openai>=1.3.0",
29-
"mcp>=1.10.1",
29+
"mcp<1.23.4,>=1.10.1",
3030
"aiohttp",
3131
"httpx>=0.27.0",
3232
"httpx-sse>=0.4.0",

tests/tools/mcp_tool/test_mcp_toolset.py

Lines changed: 197 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,9 +4,11 @@
44
#
55
# tRPC-Agent-Python is licensed under Apache-2.0.
66

7+
import asyncio
78
from unittest.mock import AsyncMock, MagicMock, patch
89

910
import pytest
11+
from mcp import types as mcp_types
1012
from mcp import StdioServerParameters as McpStdioServerParameters
1113
from mcp.types import ListToolsResult, Tool as McpBaseTool
1214

@@ -26,6 +28,13 @@ def _stdio_conn():
2628
)
2729

2830

31+
def _server_capabilities(list_changed: bool | None = None):
32+
tools_capability = None
33+
if list_changed is not None:
34+
tools_capability = mcp_types.ToolsCapability(listChanged=list_changed)
35+
return mcp_types.ServerCapabilities(tools=tools_capability)
36+
37+
2938
# ---------------------------------------------------------------------------
3039
# Tests: __init__
3140
# ---------------------------------------------------------------------------
@@ -70,6 +79,15 @@ def test_session_group_params_custom(self):
7079
ts = MCPToolset(connection_params=_stdio_conn(), session_group_params={"key": "val"})
7180
assert ts._session_group_params == {"key": "val"}
7281

82+
def test_tools_cache_enabled_by_default(self):
83+
ts = MCPToolset(connection_params=_stdio_conn())
84+
assert ts._cache_tools is True
85+
assert ts._tools_cache_ttl == 60.0
86+
87+
def test_rejects_negative_tools_cache_ttl(self):
88+
with pytest.raises(ValueError, match="tools_cache_ttl must be non-negative"):
89+
MCPToolset(connection_params=_stdio_conn(), tools_cache_ttl=-1)
90+
7391

7492
# ---------------------------------------------------------------------------
7593
# Tests: _checker_required_params
@@ -294,6 +312,185 @@ async def test_get_tools_with_custom_mcp_tool_cls(self):
294312
custom_cls.assert_called_once()
295313
assert len(tools) == 1
296314

315+
@pytest.mark.asyncio
316+
async def test_get_tools_reuses_cached_list_tools_response(self):
317+
ts = MCPToolset(connection_params=_stdio_conn())
318+
319+
mock_mgr = MagicMock(spec=MCPSessionManager)
320+
mock_session = AsyncMock()
321+
mock_mgr.create_session = AsyncMock(return_value=mock_session)
322+
323+
mcp_tools = [
324+
McpBaseTool(name="tool_a", description="desc_a", inputSchema={"type": "object"}),
325+
]
326+
mock_session.list_tools = AsyncMock(return_value=ListToolsResult(tools=mcp_tools))
327+
328+
with patch.object(ts, "initialize"):
329+
ts._mcp_session_manager = mock_mgr
330+
first = await ts.get_tools()
331+
second = await ts.get_tools()
332+
333+
assert [tool.name for tool in first] == ["tool_a"]
334+
assert [tool.name for tool in second] == ["tool_a"]
335+
mock_session.list_tools.assert_awaited_once()
336+
337+
@pytest.mark.asyncio
338+
async def test_get_tools_can_disable_tools_cache(self):
339+
ts = MCPToolset(connection_params=_stdio_conn(), cache_tools=False)
340+
341+
mock_mgr = MagicMock(spec=MCPSessionManager)
342+
mock_session = AsyncMock()
343+
mock_mgr.create_session = AsyncMock(return_value=mock_session)
344+
345+
mock_session.list_tools = AsyncMock(
346+
return_value=ListToolsResult(
347+
tools=[
348+
McpBaseTool(name="tool_a", description="desc_a", inputSchema={"type": "object"}),
349+
]
350+
))
351+
352+
with patch.object(ts, "initialize"):
353+
ts._mcp_session_manager = mock_mgr
354+
await ts.get_tools()
355+
await ts.get_tools()
356+
357+
assert mock_session.list_tools.await_count == 2
358+
359+
@pytest.mark.asyncio
360+
async def test_clear_tools_cache_forces_refresh(self):
361+
ts = MCPToolset(connection_params=_stdio_conn())
362+
363+
mock_mgr = MagicMock(spec=MCPSessionManager)
364+
mock_session = AsyncMock()
365+
mock_mgr.create_session = AsyncMock(return_value=mock_session)
366+
367+
mock_session.list_tools = AsyncMock(
368+
side_effect=[
369+
ListToolsResult(
370+
tools=[
371+
McpBaseTool(name="tool_a", description="desc_a", inputSchema={"type": "object"}),
372+
]),
373+
ListToolsResult(
374+
tools=[
375+
McpBaseTool(name="tool_b", description="desc_b", inputSchema={"type": "object"}),
376+
]),
377+
])
378+
379+
with patch.object(ts, "initialize"):
380+
ts._mcp_session_manager = mock_mgr
381+
first = await ts.get_tools()
382+
ts.clear_tools_cache()
383+
second = await ts.get_tools()
384+
385+
assert [tool.name for tool in first] == ["tool_a"]
386+
assert [tool.name for tool in second] == ["tool_b"]
387+
assert mock_session.list_tools.await_count == 2
388+
389+
@pytest.mark.asyncio
390+
async def test_tools_cache_ttl_expires(self):
391+
ts = MCPToolset(connection_params=_stdio_conn(), tools_cache_ttl=1)
392+
393+
mock_mgr = MagicMock(spec=MCPSessionManager)
394+
mock_session = AsyncMock()
395+
mock_mgr.create_session = AsyncMock(return_value=mock_session)
396+
397+
mock_session.list_tools = AsyncMock(
398+
side_effect=[
399+
ListToolsResult(
400+
tools=[
401+
McpBaseTool(name="tool_a", description="desc_a", inputSchema={"type": "object"}),
402+
]),
403+
ListToolsResult(
404+
tools=[
405+
McpBaseTool(name="tool_b", description="desc_b", inputSchema={"type": "object"}),
406+
]),
407+
])
408+
409+
with patch.object(ts, "initialize"), patch(
410+
"trpc_agent_sdk.tools.mcp_tool._mcp_toolset.time.monotonic",
411+
side_effect=[100.0, 100.5, 101.1, 101.1, 101.1],
412+
):
413+
ts._mcp_session_manager = mock_mgr
414+
first = await ts.get_tools()
415+
cached = await ts.get_tools()
416+
refreshed = await ts.get_tools()
417+
418+
assert [tool.name for tool in first] == ["tool_a"]
419+
assert [tool.name for tool in cached] == ["tool_a"]
420+
assert [tool.name for tool in refreshed] == ["tool_b"]
421+
assert mock_session.list_tools.await_count == 2
422+
423+
@pytest.mark.asyncio
424+
async def test_list_changed_capability_uses_notification_driven_cache(self):
425+
ts = MCPToolset(connection_params=_stdio_conn(), tools_cache_ttl=1)
426+
427+
mock_mgr = MagicMock(spec=MCPSessionManager)
428+
mock_session = AsyncMock()
429+
mock_session.get_server_capabilities = MagicMock(return_value=_server_capabilities(list_changed=True))
430+
mock_mgr.create_session = AsyncMock(return_value=mock_session)
431+
432+
mock_session.list_tools = AsyncMock(
433+
return_value=ListToolsResult(
434+
tools=[
435+
McpBaseTool(name="tool_a", description="desc_a", inputSchema={"type": "object"}),
436+
]
437+
))
438+
439+
with patch.object(ts, "initialize"), patch(
440+
"trpc_agent_sdk.tools.mcp_tool._mcp_toolset.time.monotonic",
441+
return_value=100.0,
442+
):
443+
ts._mcp_session_manager = mock_mgr
444+
first = await ts.get_tools()
445+
second = await ts.get_tools()
446+
447+
assert [tool.name for tool in first] == ["tool_a"]
448+
assert [tool.name for tool in second] == ["tool_a"]
449+
mock_session.list_tools.assert_awaited_once()
450+
451+
@pytest.mark.asyncio
452+
async def test_tool_list_changed_notification_clears_cache_and_chains_handler(self):
453+
user_message_handler = AsyncMock()
454+
ts = MCPToolset(
455+
connection_params=_stdio_conn(),
456+
session_group_params={"message_handler": user_message_handler},
457+
)
458+
ts._tools_cache = ListToolsResult(
459+
tools=[
460+
McpBaseTool(name="tool_a", description="desc_a", inputSchema={"type": "object"}),
461+
])
462+
ts._tools_cache_updated_at = 100.0
463+
464+
params = ts._build_session_group_params()
465+
notification = mcp_types.ServerNotification(mcp_types.ToolListChangedNotification())
466+
await params["message_handler"](notification)
467+
468+
assert ts._tools_cache is None
469+
assert ts._tools_cache_updated_at is None
470+
user_message_handler.assert_awaited_once_with(notification)
471+
472+
@pytest.mark.asyncio
473+
async def test_concurrent_get_tools_shares_cache_fill(self):
474+
ts = MCPToolset(connection_params=_stdio_conn())
475+
476+
mock_mgr = MagicMock(spec=MCPSessionManager)
477+
mock_session = AsyncMock()
478+
mock_mgr.create_session = AsyncMock(return_value=mock_session)
479+
mock_session.list_tools = AsyncMock(
480+
return_value=ListToolsResult(
481+
tools=[
482+
McpBaseTool(name="tool_a", description="desc_a", inputSchema={"type": "object"}),
483+
]
484+
))
485+
486+
with patch.object(ts, "initialize"):
487+
ts._mcp_session_manager = mock_mgr
488+
first, second = await asyncio.gather(ts.get_tools(), ts.get_tools())
489+
490+
assert [tool.name for tool in first] == ["tool_a"]
491+
assert [tool.name for tool in second] == ["tool_a"]
492+
mock_session.list_tools.assert_awaited_once()
493+
297494

298495
# ---------------------------------------------------------------------------
299496
# Tests: close

0 commit comments

Comments
 (0)