From 2b1679ef086729f1fa2558e8a4e6b723261ce8b2 Mon Sep 17 00:00:00 2001 From: QuentinBisson Date: Fri, 17 Jul 2026 16:15:19 +0200 Subject: [PATCH] feat(tools): support elicitation_callback in McpToolset Thread an optional elicitation_callback through McpToolset, MCPSessionManager, and SessionContext into the underlying ClientSession, mirroring the existing sampling_callback plumbing. Providing a callback makes the MCP client declare the elicitation capability, so servers can use elicitation/create (including URL-mode elicitation per SEP-1036) for out-of-band flows such as auth challenges instead of failing opaquely inside the toolset. --- .../adk/tools/mcp_tool/mcp_session_manager.py | 7 ++++ src/google/adk/tools/mcp_tool/mcp_toolset.py | 9 +++++ .../adk/tools/mcp_tool/session_context.py | 7 ++++ .../mcp_tool/test_mcp_session_manager.py | 37 +++++++++++++++++++ .../tools/mcp_tool/test_mcp_toolset.py | 28 ++++++++++++++ .../tools/mcp_tool/test_session_context.py | 30 +++++++++++++++ 6 files changed, 118 insertions(+) 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 3a61929e76d..98e037c5044 100644 --- a/src/google/adk/tools/mcp_tool/mcp_session_manager.py +++ b/src/google/adk/tools/mcp_tool/mcp_session_manager.py @@ -59,6 +59,7 @@ class AsyncAuthorizedSession: # pylint: disable=g-bad-classes from mcp import ClientSession from mcp import SamplingCapability from mcp import StdioServerParameters +from mcp.client.session import ElicitationFnT from mcp.client.session import SamplingFnT from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client @@ -536,6 +537,7 @@ def __init__( *, sampling_callback: SamplingFnT | None = None, sampling_capabilities: SamplingCapability | None = None, + elicitation_callback: ElicitationFnT | None = None, ): """Initializes the MCP session manager. @@ -548,9 +550,13 @@ def __init__( sampling_callback: Optional callback to handle sampling requests from the MCP server. sampling_capabilities: Optional capabilities for sampling. + elicitation_callback: Optional callback to handle elicitation requests + from the MCP server (``elicitation/create``), including URL-mode + elicitations used for out-of-band flows such as auth challenges. """ self._sampling_callback = sampling_callback self._sampling_capabilities = sampling_capabilities + self._elicitation_callback = elicitation_callback if isinstance(connection_params, StdioServerParameters): # So far timeout is not configurable. Given MCP is still evolving, we @@ -990,6 +996,7 @@ async def create_session( is_stdio=is_stdio, sampling_callback=self._sampling_callback, sampling_capabilities=self._sampling_capabilities, + elicitation_callback=self._elicitation_callback, ) if is_feature_enabled(FeatureName._MCP_GRACEFUL_ERROR_HANDLING): # pylint: disable=protected-access diff --git a/src/google/adk/tools/mcp_tool/mcp_toolset.py b/src/google/adk/tools/mcp_tool/mcp_toolset.py index b9c210735ac..df45cbe133a 100644 --- a/src/google/adk/tools/mcp_tool/mcp_toolset.py +++ b/src/google/adk/tools/mcp_tool/mcp_toolset.py @@ -32,6 +32,7 @@ from mcp import SamplingCapability from mcp import StdioServerParameters +from mcp.client.session import ElicitationFnT from mcp.client.session import SamplingFnT from mcp.shared.session import ProgressFnT from mcp.types import ListResourcesResult @@ -121,6 +122,7 @@ def __init__( use_mcp_resources: Optional[bool] = False, sampling_callback: Optional[SamplingFnT] = None, sampling_capabilities: Optional[SamplingCapability] = None, + elicitation_callback: Optional[ElicitationFnT] = None, credential_key: str | None = None, ): """Initializes the McpToolset. @@ -161,6 +163,11 @@ def __init__( sampling_callback: Optional callback to handle sampling requests from the MCP server. sampling_capabilities: Optional capabilities for sampling. + elicitation_callback: Optional callback to handle elicitation requests + from the MCP server (``elicitation/create``), including URL-mode + elicitations used for out-of-band flows such as auth challenges. + Providing a callback makes the client declare the elicitation + capability during initialization. credential_key: A user specified key used to load and save this credential in a credential service. Used with auth_scheme. """ @@ -169,6 +176,7 @@ def __init__( self._sampling_callback = sampling_callback self._sampling_capabilities = sampling_capabilities + self._elicitation_callback = elicitation_callback if not connection_params: raise ValueError("Missing connection params in McpToolset.") @@ -184,6 +192,7 @@ def __init__( errlog=self._errlog, sampling_callback=self._sampling_callback, sampling_capabilities=self._sampling_capabilities, + elicitation_callback=self._elicitation_callback, ) self._auth_scheme = auth_scheme self._auth_credential = auth_credential diff --git a/src/google/adk/tools/mcp_tool/session_context.py b/src/google/adk/tools/mcp_tool/session_context.py index db367e8c701..e29fd7f8714 100644 --- a/src/google/adk/tools/mcp_tool/session_context.py +++ b/src/google/adk/tools/mcp_tool/session_context.py @@ -27,6 +27,7 @@ from mcp import ClientSession from mcp import SamplingCapability +from mcp.client.session import ElicitationFnT from mcp.client.session import SamplingFnT from ...features import FeatureName @@ -96,6 +97,7 @@ def __init__( *, sampling_callback: Optional[SamplingFnT] = None, sampling_capabilities: Optional[SamplingCapability] = None, + elicitation_callback: Optional[ElicitationFnT] = None, ): """ Args: @@ -108,6 +110,8 @@ def __init__( sampling_callback: Optional callback to handle sampling requests from the MCP server. sampling_capabilities: Optional capabilities for sampling. + elicitation_callback: Optional callback to handle elicitation requests + from the MCP server (``elicitation/create``). """ self._client = client self._timeout = timeout @@ -120,6 +124,7 @@ def __init__( self._task_lock = asyncio.Lock() self._sampling_callback = sampling_callback self._sampling_capabilities = sampling_capabilities + self._elicitation_callback = elicitation_callback @property def session(self) -> Optional[ClientSession]: @@ -320,6 +325,7 @@ async def _run(self) -> None: else None, sampling_callback=self._sampling_callback, sampling_capabilities=self._sampling_capabilities, + elicitation_callback=self._elicitation_callback, ) ) else: @@ -333,6 +339,7 @@ async def _run(self) -> None: else None, sampling_callback=self._sampling_callback, sampling_capabilities=self._sampling_capabilities, + elicitation_callback=self._elicitation_callback, ) ) # pylint: disable-next=protected-access 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 2f6a11305d5..fe42fab8a1d 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_session_manager.py @@ -406,6 +406,43 @@ async def test_create_session_stdio_new(self): # Verify enter_async_context was called (which internally calls __aenter__) mock_exit_stack.enter_async_context.assert_called_once() + @pytest.mark.asyncio + async def test_create_session_passes_elicitation_callback(self): + """Elicitation callback is forwarded to the SessionContext.""" + + async def elicitation_callback(context, params): + return {"action": "decline"} + + manager = MCPSessionManager( + self.mock_stdio_connection_params, + elicitation_callback=elicitation_callback, + ) + + mock_exit_stack = MockAsyncExitStack() + + with patch( + "google.adk.tools.mcp_tool.mcp_session_manager.stdio_client" + ) as mock_stdio: + with patch( + "google.adk.tools.mcp_tool.mcp_session_manager.AsyncExitStack" + ) as mock_exit_stack_class: + with patch( + "google.adk.tools.mcp_tool.mcp_session_manager.SessionContext" + ) as mock_session_context_class: + mock_exit_stack_class.return_value = mock_exit_stack + mock_stdio.return_value = AsyncMock() + + mock_session = AsyncMock() + mock_session_context = MockSessionContext(session=mock_session) + mock_session_context_class.return_value = mock_session_context + mock_exit_stack.enter_async_context.return_value = mock_session + + await manager.create_session() + + mock_session_context_class.assert_called_once() + _, kwargs = mock_session_context_class.call_args + assert kwargs["elicitation_callback"] is elicitation_callback + @pytest.mark.asyncio async def test_create_session_reuse_existing(self): """Test reusing an existing connected session.""" diff --git a/tests/unittests/tools/mcp_tool/test_mcp_toolset.py b/tests/unittests/tools/mcp_tool/test_mcp_toolset.py index fd4b5fe621b..59b64eda5d3 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_toolset.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_toolset.py @@ -788,6 +788,34 @@ async def mock_sampling_handler(messages, params=None, context=None): assert result["role"] == "assistant" assert result["content"]["text"] == "sampling response" + @pytest.mark.asyncio + async def test_elicitation_callback_plumbed_to_session_manager(self): + """Elicitation callback reaches the session manager unchanged.""" + + async def mock_elicitation_handler(context, params): + return {"action": "decline"} + + toolset = McpToolset( + connection_params=StreamableHTTPConnectionParams( + url="http://localhost:9999", + timeout=10, + ), + elicitation_callback=mock_elicitation_handler, + ) + + assert toolset._elicitation_callback is mock_elicitation_handler + assert ( + toolset._mcp_session_manager._elicitation_callback + is mock_elicitation_handler + ) + + @pytest.mark.asyncio + async def test_elicitation_callback_defaults_to_none(self): + toolset = McpToolset(connection_params=self.mock_stdio_params) + + assert toolset._elicitation_callback is None + assert toolset._mcp_session_manager._elicitation_callback is None + @pytest.mark.asyncio async def test_get_auth_headers_includes_additional_headers(self): credential = AuthCredential( diff --git a/tests/unittests/tools/mcp_tool/test_session_context.py b/tests/unittests/tools/mcp_tool/test_session_context.py index 9634a4013a2..fa139f42ab6 100644 --- a/tests/unittests/tools/mcp_tool/test_session_context.py +++ b/tests/unittests/tools/mcp_tool/test_session_context.py @@ -116,6 +116,36 @@ async def test_start_success_ready_event_set_and_session_returned(self): # Clean up await session_context.close() + @pytest.mark.asyncio + async def test_elicitation_callback_passed_to_client_session(self): + """Elicitation callback is forwarded to the ClientSession.""" + + async def elicitation_callback(context, params): + return {'action': 'decline'} + + mock_client = MockClient() + session_context = SessionContext( + mock_client, + timeout=5.0, + sse_read_timeout=None, + elicitation_callback=elicitation_callback, + ) + + mock_session = MockClientSession() + + with patch( + 'google.adk.tools.mcp_tool.session_context.ClientSession' + ) as mock_session_class: + mock_session_class.return_value = mock_session + + await session_context.start() + + mock_session_class.assert_called_once() + _, kwargs = mock_session_class.call_args + assert kwargs['elicitation_callback'] is elicitation_callback + + await session_context.close() + @pytest.mark.asyncio async def test_start_raises_connection_error_on_exception(self): """Test that start() raises ConnectionError when exception occurs."""