44#
55# tRPC-Agent-Python is licensed under Apache-2.0.
66
7+ import asyncio
78from unittest .mock import AsyncMock , MagicMock , patch
89
910import pytest
11+ from mcp import types as mcp_types
1012from mcp import StdioServerParameters as McpStdioServerParameters
1113from 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