diff --git a/src/google/adk/tools/mcp_tool/mcp_toolset.py b/src/google/adk/tools/mcp_tool/mcp_toolset.py index 35a43754..30228813 100644 --- a/src/google/adk/tools/mcp_tool/mcp_toolset.py +++ b/src/google/adk/tools/mcp_tool/mcp_toolset.py @@ -215,37 +215,25 @@ class McpToolset(BaseToolset): async def read_resource( self, name: str, readonly_context: Optional[ReadonlyContext] = None ) -> Any: - """Fetches and returns the content of the named resource. - - This method will handle content decoding based on the MIME type reported by - the MCP server (e.g., JSON, text, base64 for binary). + """Fetches and returns a list of contents of the named resource. Args: name: The name of the resource to fetch. readonly_context: Context used to provide headers for the MCP session. Returns: - The content of the resource, decoded based on MIME type and encoding. + List of contents of the resource. """ + resource_info = await self.get_resource_info(name, readonly_context) + if "uri" not in resource_info: + raise ValueError(f"Resource '{name}' has no URI.") + result: Any = await self._execute_with_session( - lambda session: session.get_resource(name=name), + lambda session: session.read_resource(uri=resource_info["uri"]), f"Failed to get resource {name} from MCP server", readonly_context, ) - - content = result.content - if result.encoding == "base64": - decoded_bytes = base64.b64decode(content) - if result.resource.mime_type == "application/json": - return json.loads(decoded_bytes.decode("utf-8")) - if result.resource.mime_type.startswith("text/"): - return decoded_bytes.decode("utf-8") - return decoded_bytes # Return as bytes for other binary types - - if result.resource.mime_type == "application/json": - return json.loads(content) - - return content + return result.contents async def list_resources( self, readonly_context: Optional[ReadonlyContext] = None diff --git a/tests/unittests/tools/mcp_tool/test_mcp_toolset.py b/tests/unittests/tools/mcp_tool/test_mcp_toolset.py index 342f2e77..2c925609 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_toolset.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_toolset.py @@ -32,8 +32,11 @@ from google.adk.tools.mcp_tool.mcp_session_manager import StreamableHTTPConnecti from google.adk.tools.mcp_tool.mcp_tool import MCPTool from google.adk.tools.mcp_tool.mcp_toolset import McpToolset from mcp import StdioServerParameters +from mcp.types import BlobResourceContents from mcp.types import ListResourcesResult +from mcp.types import ReadResourceResult from mcp.types import Resource +from mcp.types import TextResourceContents import pytest @@ -435,56 +438,66 @@ class TestMcpToolset: await toolset.get_resource_info("other.json") @pytest.mark.parametrize( - "name,mime_type,content,encoding,expected_result", + "name,mime_type,content,encoding", [ - ("file1.txt", "text/plain", "hello world", None, "hello world"), + ("file1.txt", "text/plain", "hello world", None), ( "data.json", "application/json", '{"key": "value"}', None, - {"key": "value"}, ), ( "file1_b64.txt", "text/plain", base64.b64encode(b"hello world").decode("ascii"), "base64", - "hello world", ), ( "data_b64.json", "application/json", base64.b64encode(b'{"key": "value"}').decode("ascii"), "base64", - {"key": "value"}, ), ( "data.bin", "application/octet-stream", base64.b64encode(b"\x01\x02\x03").decode("ascii"), "base64", - b"\x01\x02\x03", ), ], ) @pytest.mark.asyncio - async def test_read_resource( - self, name, mime_type, content, encoding, expected_result - ): + async def test_read_resource(self, name, mime_type, content, encoding): """Test reading various resource types.""" - get_resource_result = MagicMock() - get_resource_result.resource = Resource( - name=name, mime_type=mime_type, uri=f"file:///{name}" + uri = f"file:///{name}" + # Mock list_resources for get_resource_info + resources = [Resource(name=name, mime_type=mime_type, uri=uri)] + list_resources_result = ListResourcesResult(resources=resources) + self.mock_session.list_resources = AsyncMock( + return_value=list_resources_result + ) + + # Mock read_resource + if encoding == "base64": + contents = [ + BlobResourceContents(uri=uri, mimeType=mime_type, blob=content) + ] + else: + contents = [ + TextResourceContents(uri=uri, mimeType=mime_type, text=content) + ] + + read_resource_result = ReadResourceResult(contents=contents) + self.mock_session.read_resource = AsyncMock( + return_value=read_resource_result ) - get_resource_result.content = content - get_resource_result.encoding = encoding - self.mock_session.get_resource = AsyncMock(return_value=get_resource_result) toolset = McpToolset(connection_params=self.mock_stdio_params) toolset._mcp_session_manager = self.mock_session_manager result = await toolset.read_resource(name) - assert result == expected_result - self.mock_session.get_resource.assert_called_once_with(name=name) + assert result == contents + self.mock_session.list_resources.assert_called_once() + self.mock_session.read_resource.assert_called_once_with(uri=uri)