diff --git a/src/google/adk/tools/mcp_tool/mcp_session_manager.py b/src/google/adk/tools/mcp_tool/mcp_session_manager.py index af255d94..d95d48f2 100644 --- a/src/google/adk/tools/mcp_tool/mcp_session_manager.py +++ b/src/google/adk/tools/mcp_tool/mcp_session_manager.py @@ -37,7 +37,6 @@ try: from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client from mcp.client.streamable_http import streamablehttp_client - from mcp.types import EmptyResult except ImportError as e: if sys.version_info < (3, 10): @@ -242,7 +241,7 @@ class MCPSessionManager: return base_headers - async def _is_session_disconnected(self, session: ClientSession) -> bool: + def _is_session_disconnected(self, session: ClientSession) -> bool: """Checks if a session is disconnected or closed. Args: @@ -251,24 +250,7 @@ class MCPSessionManager: Returns: True if the session is disconnected, False otherwise. """ - if session._read_stream._closed or session._write_stream._closed: - return True - - try: - response = await asyncio.wait_for(session.send_ping(), timeout=5.0) - if not isinstance(response, EmptyResult): - logger.info( - 'Session ping returns illegal response %s, treating as' - ' disconnected', - response, - ) - return True - return False - except Exception as e: - logger.info( - 'Session ping failed with error %s, treating as disconnected', e - ) - return True + return session._read_stream._closed or session._write_stream._closed def _create_client(self, merged_headers: Optional[Dict[str, str]] = None): """Creates an MCP client based on the connection parameters. @@ -343,7 +325,7 @@ class MCPSessionManager: session, exit_stack = self._sessions[session_key] # Check if the existing session is still connected - if not await self._is_session_disconnected(session): + if not self._is_session_disconnected(session): # Session is still good, return it return session else: diff --git a/src/google/adk/tools/mcp_tool/mcp_toolset.py b/src/google/adk/tools/mcp_tool/mcp_toolset.py index 429d63ab..daa88f90 100644 --- a/src/google/adk/tools/mcp_tool/mcp_toolset.py +++ b/src/google/adk/tools/mcp_tool/mcp_toolset.py @@ -175,10 +175,7 @@ class McpToolset(BaseToolset): else None ) # Get session from session manager - try: - session = await self._mcp_session_manager.create_session(headers=headers) - except Exception as e: - raise ConnectionError(f"Failed to create MCP session") from e + session = await self._mcp_session_manager.create_session(headers=headers) # Fetch available tools from the MCP server timeout_in_seconds = ( diff --git a/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py b/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py index 8eb743eb..6c001ccf 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py @@ -54,7 +54,6 @@ except ImportError as e: # Import real MCP classes try: from mcp import StdioServerParameters - from mcp.types import EmptyResult except ImportError: # Create a mock if MCP is not available class StdioServerParameters: @@ -63,9 +62,6 @@ except ImportError: self.command = command self.args = args or [] - class EmptyResult: - pass - class MockClientSession: """Mock ClientSession for testing.""" @@ -76,7 +72,6 @@ class MockClientSession: self._read_stream._closed = False self._write_stream._closed = False self.initialize = AsyncMock() - self.send_ping = AsyncMock() class MockAsyncExitStack: @@ -211,52 +206,19 @@ class TestMCPSessionManager: } assert merged == expected - @pytest.mark.asyncio - async def test_is_session_disconnected_when_connected(self): - """Test session disconnection detection when session is connected.""" + def test_is_session_disconnected(self): + """Test session disconnection detection.""" manager = MCPSessionManager(self.mock_stdio_connection_params) - session = MockClientSession() - session.send_ping.return_value = EmptyResult() - assert not await manager._is_session_disconnected(session) - session.send_ping.assert_called_once() - @pytest.mark.asyncio - async def test_is_session_disconnected_read_stream_closed(self): - """Test session disconnection detection when read stream is closed.""" - manager = MCPSessionManager(self.mock_stdio_connection_params) + # Create mock session session = MockClientSession() - session.send_ping.return_value = EmptyResult() + + # Not disconnected + assert not manager._is_session_disconnected(session) + + # Disconnected - read stream closed session._read_stream._closed = True - assert await manager._is_session_disconnected(session) - session.send_ping.assert_not_called() - - @pytest.mark.asyncio - async def test_is_session_disconnected_write_stream_closed(self): - """Test session disconnection detection when write stream is closed.""" - manager = MCPSessionManager(self.mock_stdio_connection_params) - session = MockClientSession() - session.send_ping.return_value = EmptyResult() - session._write_stream._closed = True - assert await manager._is_session_disconnected(session) - session.send_ping.assert_not_called() - - @pytest.mark.asyncio - async def test_is_session_disconnected_ping_fails(self): - """Test session disconnection detection when ping fails.""" - manager = MCPSessionManager(self.mock_stdio_connection_params) - session = MockClientSession() - session.send_ping.side_effect = Exception("Ping failed") - assert await manager._is_session_disconnected(session) - session.send_ping.assert_called_once() - - @pytest.mark.asyncio - async def test_is_session_disconnected_ping_returns_wrong_result(self): - """Test session disconnection detection when ping returns wrong result.""" - manager = MCPSessionManager(self.mock_stdio_connection_params) - session = MockClientSession() - session.send_ping.return_value = "Wrong result" - assert await manager._is_session_disconnected(session) - session.send_ping.assert_called_once() + assert manager._is_session_disconnected(session) @pytest.mark.asyncio async def test_create_session_stdio_new(self): @@ -309,7 +271,6 @@ class TestMCPSessionManager: # Session is connected existing_session._read_stream._closed = False existing_session._write_stream._closed = False - existing_session.send_ping.return_value = EmptyResult() session = await manager.create_session()