diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 039ac428352..cb9f11f82de 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -86,6 +86,10 @@ class _LenientListPromptsResult(types.ListPromptsResult): prompts: List[Prompt] = Field(default_factory=list) +class _LenientListToolsResult(types.ListToolsResult): + tools: List[MCPTool] = Field(default_factory=list) + + class _LenientListResourcesResult(types.ListResourcesResult): resources: List[Resource] = Field(default_factory=list) @@ -521,7 +525,10 @@ class MCPClient: ) async def _list_tools_operation(session: ClientSession): - return await session.list_tools() + return await session.send_request( + types.ClientRequest(types.ListToolsRequest(params=None)), + _LenientListToolsResult, + ) try: result = await self.run_with_session(_list_tools_operation) diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 468fe61d2e5..c873256cda8 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -425,6 +425,53 @@ class TestMCPClientInstructionsCapture: assert client._last_initialize_instructions is None +# --------------------------------------------------------------------------- +# MCP list result tolerance +# --------------------------------------------------------------------------- + + +class TestMCPListOperations: + """MCP list operations should tolerate omitted empty list fields.""" + + async def _run_list_operation(self, method_name, result_type): + client = MCPClient(server_url="http://example.com/mcp", transport_type="http") + session = AsyncMock() + + async def send_request(_request, actual_result_type): + assert actual_result_type is result_type + return actual_result_type.model_validate({}) + + session.send_request = AsyncMock(side_effect=send_request) + + async def run_with_session(operation): + return await operation(session) + + client.run_with_session = AsyncMock(side_effect=run_with_session) + + result = await getattr(client, method_name)() + + assert result == [] + session.send_request.assert_awaited_once() + + @pytest.mark.asyncio + async def test_list_tools_tolerates_missing_tools_field(self): + await self._run_list_operation( + "list_tools", mcp_client_module._LenientListToolsResult + ) + + @pytest.mark.asyncio + async def test_list_resources_tolerates_missing_resources_field(self): + await self._run_list_operation( + "list_resources", mcp_client_module._LenientListResourcesResult + ) + + @pytest.mark.asyncio + async def test_list_prompts_tolerates_missing_prompts_field(self): + await self._run_list_operation( + "list_prompts", mcp_client_module._LenientListPromptsResult + ) + + # --------------------------------------------------------------------------- # Transport error surfacing # --------------------------------------------------------------------------- @@ -542,44 +589,6 @@ class TestExecuteSessionOperationSurfacesTransportError: result = await client._execute_session_operation(transport_ctx, _op) assert result == "done" - @pytest.mark.asyncio - async def test_list_resources_tolerates_missing_resources_field(self): - client = MCPClient(server_url="http://example.com/mcp", transport_type="http") - session = AsyncMock() - - async def send_request(_request, result_type): - assert result_type is mcp_client_module._LenientListResourcesResult - return result_type.model_validate({}) - - session.send_request = AsyncMock(side_effect=send_request) - - async def run_with_session(operation): - return await operation(session) - - client.run_with_session = AsyncMock(side_effect=run_with_session) - - assert await client.list_resources() == [] - session.send_request.assert_awaited_once() - - @pytest.mark.asyncio - async def test_list_prompts_tolerates_missing_prompts_field(self): - client = MCPClient(server_url="http://example.com/mcp", transport_type="http") - session = AsyncMock() - - async def send_request(_request, result_type): - assert result_type is mcp_client_module._LenientListPromptsResult - return result_type.model_validate({}) - - session.send_request = AsyncMock(side_effect=send_request) - - async def run_with_session(operation): - return await operation(session) - - client.run_with_session = AsyncMock(side_effect=run_with_session) - - assert await client.list_prompts() == [] - session.send_request.assert_awaited_once() - if __name__ == "__main__": pytest.main([__file__])