Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,7 @@ optional-dependencies.all = [
"google-cloud-spanner>=3.56,<4",
"google-cloud-speech>=2.30,<3",
"google-cloud-storage>=2.18,<4",
"mcp>=1.24,<2",
"mcp>=2.0.0,<3",
"opentelemetry-exporter-gcp-logging>=1.9.0a0,<=1.12.0a0",
"opentelemetry-exporter-gcp-monitoring>=1.9.0a0,<2",
"opentelemetry-exporter-gcp-trace>=1.9,<2",
Expand Down Expand Up @@ -195,7 +195,7 @@ optional-dependencies.gcp = [
]
optional-dependencies.mcp = [
"anyio>=4.9,<5",
"mcp>=1.24,<2",
"mcp>=2.0.0,<3",
]
optional-dependencies.oci = [
"oci>=2.126", # OCI Generative AI native SDK (OCIGenAILlm)
Expand Down Expand Up @@ -243,7 +243,7 @@ optional-dependencies.test = [
"litellm>=1.84",
"llama-index-readers-file>=0.4",
"lxml>=5.3",
"mcp>=1.24,<2",
"mcp>=2.0.0,<3",
"openai>=2.20,<3",
"opentelemetry-exporter-gcp-logging>=1.9.0a0,<=1.12.0a0",
"opentelemetry-exporter-gcp-monitoring>=1.9.0a0,<2",
Expand Down
39 changes: 29 additions & 10 deletions src/google/adk/tools/mcp_tool/_agent_to_mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,8 +23,8 @@

from google.genai import types
from mcp import types as mcp_types
from mcp.server.fastmcp import Context
from mcp.server.fastmcp import FastMCP
from mcp.server.mcpserver import Context
from mcp.server.mcpserver import MCPServer

from ...agents.base_agent import BaseAgent
from ...artifacts.in_memory_artifact_service import InMemoryArtifactService
Expand Down Expand Up @@ -68,18 +68,37 @@ def _part_to_content(part: types.Part) -> Optional[mcp_types.ContentBlock]:
data = base64.b64encode(blob.data).decode("ascii")
mime = blob.mime_type or "application/octet-stream"
if mime.startswith("image/"):
return mcp_types.ImageContent(type="image", data=data, mimeType=mime)
return mcp_types.ImageContent(type="image", data=data, mime_type=mime)
if mime.startswith("audio/"):
return mcp_types.AudioContent(type="audio", data=data, mimeType=mime)
return mcp_types.AudioContent(type="audio", data=data, mime_type=mime)
return mcp_types.EmbeddedResource(
type="resource",
resource=mcp_types.BlobResourceContents(
uri=_INLINE_RESOURCE_URI, blob=data, mimeType=mime
uri=_INLINE_RESOURCE_URI, blob=data, mime_type=mime
),
)
return None


def _connection_key(ctx: Context) -> object:
"""Returns a stable per-connection key for an MCP tool call context.

In mcp 2.0, ``ctx.session`` is a new ``ServerSession`` object on every
request even over a single connection, so it can no longer key the
per-connection ADK session map. The underlying ``Connection`` object is
shared by every request on one connection, so we use it when available and
fall back to ``ctx.session`` otherwise.

Args:
ctx: The MCP tool call context.

Returns:
A hashable object that is stable across all requests on one connection.
"""
connection = getattr(ctx.session, "_connection", None)
return connection if connection is not None else ctx.session


async def _run_agent(
runner: Runner,
request: str,
Expand All @@ -106,14 +125,14 @@ async def _run_agent(
"""
session_id: Optional[str] = None
if ctx is not None and sessions is not None:
session_id = sessions.get(ctx.session)
session_id = sessions.get(_connection_key(ctx))
if session_id is None:
session = await runner.session_service.create_session(
app_name=runner.app_name, user_id=_MCP_USER_ID
)
session_id = session.id
if ctx is not None and sessions is not None:
sessions[ctx.session] = session_id
sessions[_connection_key(ctx)] = session_id
new_message = types.Content(role="user", parts=[types.Part(text=request)])
final_content: list[mcp_types.ContentBlock] = []
async for event in runner.run_async(
Expand Down Expand Up @@ -142,7 +161,7 @@ def to_mcp_server(
name: Optional[str] = None,
instructions: Optional[str] = None,
runner: Optional[Runner] = None,
) -> FastMCP:
) -> MCPServer:
"""Exposes an ADK agent as an MCP server.

The returned server registers a single MCP tool that runs the agent: an MCP
Expand All @@ -166,7 +185,7 @@ def to_mcp_server(
services.

Returns:
A ``FastMCP`` server exposing the agent as a single tool.
A ``MCPServer`` server exposing the agent as a single tool.

Example::

Expand All @@ -175,7 +194,7 @@ def to_mcp_server(
server.run(transport="stdio")
"""
tool_name = name or agent.name or "adk_agent"
server = FastMCP(name=tool_name, instructions=instructions)
server = MCPServer(name=tool_name, instructions=instructions)
agent_runner = runner if runner is not None else _build_runner(agent)
# Maps each MCP connection to its ADK session; WeakKeyDictionary drops the
# entry when the connection is garbage-collected. pylint wrongly flags the
Expand Down
2 changes: 1 addition & 1 deletion src/google/adk/tools/mcp_tool/conversion_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ def adk_to_mcp_tool_type(tool: BaseTool) -> mcp_types.Tool:
return mcp_types.Tool(
name=tool.name,
description=tool.description,
inputSchema=input_schema,
input_schema=input_schema,
)


Expand Down
18 changes: 15 additions & 3 deletions src/google/adk/tools/mcp_tool/mcp_session_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,6 @@ class AsyncAuthorizedSession: # pylint: disable=g-bad-classes
from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client
from mcp.client.streamable_http import create_mcp_http_client as _create_mcp_http_client
from mcp.client.streamable_http import McpHttpClientFactory
from mcp.client.streamable_http import streamable_http_client
from pydantic import BaseModel
from pydantic import ConfigDict
Expand Down Expand Up @@ -215,8 +214,21 @@ class SseConnectionParams(BaseModel):


@runtime_checkable
class CheckableMcpHttpClientFactory(McpHttpClientFactory, Protocol):
pass
class CheckableMcpHttpClientFactory(Protocol):
"""Factory protocol for creating custom HTTPX async clients.

In mcp 2.0 the upstream ``McpHttpClientFactory`` protocol lives in the
private ``mcp.shared._httpx_utils`` module, so we declare the equivalent
shape locally rather than depend on a private import.
"""

def __call__(
self,
headers: dict[str, str] | None = None,
timeout: httpx.Timeout | None = None,
auth: httpx.Auth | None = None,
) -> httpx.AsyncClient:
...


class _DebugHttpxClientFactory:
Expand Down
14 changes: 7 additions & 7 deletions src/google/adk/tools/mcp_tool/mcp_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,8 +28,8 @@

from fastapi.openapi.models import APIKeyIn
from google.genai.types import FunctionDeclaration
from mcp.shared.exceptions import McpError
from mcp.shared.session import ProgressFnT
from mcp.shared.dispatcher import ProgressFnT
from mcp.shared.exceptions import MCPError
from mcp.types import Tool as McpBaseTool
from opentelemetry import propagate
from typing_extensions import override
Expand Down Expand Up @@ -201,8 +201,8 @@ def _get_declaration(self) -> FunctionDeclaration:
Returns:
FunctionDeclaration: The Gemini function declaration for the tool.
"""
input_schema = self._mcp_tool.inputSchema
output_schema = self._mcp_tool.outputSchema
input_schema = self._mcp_tool.input_schema
output_schema = self._mcp_tool.output_schema
if is_feature_enabled(FeatureName.JSON_SCHEMA_FOR_FUNC_DECL):
function_decl = FunctionDeclaration(
name=self.name,
Expand Down Expand Up @@ -375,8 +375,8 @@ async def run_async(
# any AGW policy) returns a 403 mid-tool-call.
try:
return await super().run_async(args=args, tool_context=tool_context)
except McpError as e:
logger.warning("MCP tool execution failed with McpError: %s", e)
except MCPError as e:
logger.warning("MCP tool execution failed with MCPError: %s", e)
return {"error": f"MCP tool execution failed: {e}"}
except Exception as e: # pylint: disable=broad-exception-caught
logger.warning(
Expand Down Expand Up @@ -489,7 +489,7 @@ async def _run_async_impl(

def _detect_error_in_response(self, response: Any) -> str | None:
"""Telemetry hook: returns an error type if the response indicates an error."""
if isinstance(response, dict) and response.get("isError"):
if isinstance(response, dict) and response.get("is_error"):
return "MCP_TOOL_ERROR"
return None

Expand Down
2 changes: 1 addition & 1 deletion src/google/adk/tools/mcp_tool/mcp_toolset.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@
from mcp import StdioServerParameters
from mcp.client.session import ElicitationFnT
from mcp.client.session import SamplingFnT
from mcp.shared.session import ProgressFnT
from mcp.shared.dispatcher import ProgressFnT
from mcp.types import ListResourcesResult
from mcp.types import ListToolsResult
from pydantic import model_validator
Expand Down
5 changes: 2 additions & 3 deletions src/google/adk/tools/mcp_tool/session_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,6 @@
import asyncio
from contextlib import AbstractAsyncContextManager
from contextlib import AsyncExitStack
from datetime import timedelta
import logging
from types import TracebackType
from typing import Any
Expand Down Expand Up @@ -321,7 +320,7 @@ async def _run(self) -> None:
session = await exit_stack.enter_async_context(
ClientSession(
*transports[:2],
read_timeout_seconds=timedelta(seconds=self._timeout)
read_timeout_seconds=float(self._timeout)
if self._timeout is not None
else None,
sampling_callback=self._sampling_callback,
Expand All @@ -335,7 +334,7 @@ async def _run(self) -> None:
session = await exit_stack.enter_async_context(
ClientSession(
*transports[:2],
read_timeout_seconds=timedelta(seconds=self._sse_read_timeout)
read_timeout_seconds=float(self._sse_read_timeout)
if self._sse_read_timeout is not None
else None,
sampling_callback=self._sampling_callback,
Expand Down
50 changes: 44 additions & 6 deletions tests/unittests/tools/mcp_tool/test_agent_to_mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,19 +15,57 @@
from __future__ import annotations

import base64
from contextlib import asynccontextmanager
from types import SimpleNamespace
from typing import AsyncGenerator

import anyio
from google.adk.agents.base_agent import BaseAgent
from google.adk.agents.invocation_context import InvocationContext
from google.adk.events.event import Event
from google.adk.tools.mcp_tool._agent_to_mcp import _run_agent
from google.adk.tools.mcp_tool._agent_to_mcp import to_mcp_server
from google.genai import types
from mcp.shared.memory import create_connected_server_and_client_session
from mcp import ClientSession
from mcp.server.mcpserver import MCPServer
from mcp.shared.memory import create_client_server_memory_streams
import pytest


@asynccontextmanager
async def _connected_client_session(server: MCPServer):
"""Connects an in-memory ClientSession to an MCPServer for testing.

This replaces the ``create_connected_server_and_client_session`` helper that
was removed in mcp 2.0. It runs the server's low-level transport on one end
of an in-memory stream pair and yields a connected, initialized
ClientSession on the other.

Args:
server: The MCPServer to connect to.

Yields:
An initialized ClientSession connected to the server.
"""
async with create_client_server_memory_streams() as (
client_streams,
server_streams,
):
client_read, client_write = client_streams
server_read, server_write = server_streams
lowlevel_server = server._lowlevel_server # pylint: disable=protected-access
async with anyio.create_task_group() as task_group:
task_group.start_soon(
lowlevel_server.run,
server_read,
server_write,
lowlevel_server.create_initialization_options(),
)
async with ClientSession(client_read, client_write) as session:
await session.initialize()
yield session


class _EchoAgent(BaseAgent):
"""Minimal agent that emits a single final text event."""

Expand Down Expand Up @@ -111,7 +149,7 @@ async def test_to_mcp_server_registers_agent_as_single_tool():
assert len(tools) == 1
assert tools[0].name == "my_agent"
assert tools[0].description == "does useful things"
assert "request" in tools[0].inputSchema["properties"]
assert "request" in tools[0].input_schema["properties"]


@pytest.mark.asyncio
Expand All @@ -129,10 +167,10 @@ async def test_call_tool_runs_agent_end_to_end():
agent = _EchoAgent(name="assistant")
server = to_mcp_server(agent)

async with create_connected_server_and_client_session(server) as client:
async with _connected_client_session(server) as client:
result = await client.call_tool("assistant", {"request": "hi"})

assert not result.isError
assert not result.is_error
assert "hello from the agent" in result.content[0].text


Expand Down Expand Up @@ -174,7 +212,7 @@ async def test_run_agent_maps_image_output_to_image_content():

assert len(result) == 1
assert result[0].type == "image"
assert result[0].mimeType == "image/png"
assert result[0].mime_type == "image/png"
assert base64.b64decode(result[0].data) == png


Expand Down Expand Up @@ -209,7 +247,7 @@ async def test_call_tool_reuses_session_across_calls_on_one_connection():
runner = _FakeRunner([_text_event("ok")])
server = to_mcp_server(agent, runner=runner)

async with create_connected_server_and_client_session(server) as client:
async with _connected_client_session(server) as client:
await client.call_tool("assistant", {"request": "first"})
await client.call_tool("assistant", {"request": "second"})

Expand Down
Loading