Skip to content

Commit 2b1679e

Browse files
committed
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.
1 parent fd006db commit 2b1679e

6 files changed

Lines changed: 118 additions & 0 deletions

File tree

src/google/adk/tools/mcp_tool/mcp_session_manager.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -59,6 +59,7 @@ class AsyncAuthorizedSession: # pylint: disable=g-bad-classes
5959
from mcp import ClientSession
6060
from mcp import SamplingCapability
6161
from mcp import StdioServerParameters
62+
from mcp.client.session import ElicitationFnT
6263
from mcp.client.session import SamplingFnT
6364
from mcp.client.sse import sse_client
6465
from mcp.client.stdio import stdio_client
@@ -536,6 +537,7 @@ def __init__(
536537
*,
537538
sampling_callback: SamplingFnT | None = None,
538539
sampling_capabilities: SamplingCapability | None = None,
540+
elicitation_callback: ElicitationFnT | None = None,
539541
):
540542
"""Initializes the MCP session manager.
541543
@@ -548,9 +550,13 @@ def __init__(
548550
sampling_callback: Optional callback to handle sampling requests from the
549551
MCP server.
550552
sampling_capabilities: Optional capabilities for sampling.
553+
elicitation_callback: Optional callback to handle elicitation requests
554+
from the MCP server (``elicitation/create``), including URL-mode
555+
elicitations used for out-of-band flows such as auth challenges.
551556
"""
552557
self._sampling_callback = sampling_callback
553558
self._sampling_capabilities = sampling_capabilities
559+
self._elicitation_callback = elicitation_callback
554560

555561
if isinstance(connection_params, StdioServerParameters):
556562
# So far timeout is not configurable. Given MCP is still evolving, we
@@ -990,6 +996,7 @@ async def create_session(
990996
is_stdio=is_stdio,
991997
sampling_callback=self._sampling_callback,
992998
sampling_capabilities=self._sampling_capabilities,
999+
elicitation_callback=self._elicitation_callback,
9931000
)
9941001

9951002
if is_feature_enabled(FeatureName._MCP_GRACEFUL_ERROR_HANDLING): # pylint: disable=protected-access

src/google/adk/tools/mcp_tool/mcp_toolset.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@
3232

3333
from mcp import SamplingCapability
3434
from mcp import StdioServerParameters
35+
from mcp.client.session import ElicitationFnT
3536
from mcp.client.session import SamplingFnT
3637
from mcp.shared.session import ProgressFnT
3738
from mcp.types import ListResourcesResult
@@ -121,6 +122,7 @@ def __init__(
121122
use_mcp_resources: Optional[bool] = False,
122123
sampling_callback: Optional[SamplingFnT] = None,
123124
sampling_capabilities: Optional[SamplingCapability] = None,
125+
elicitation_callback: Optional[ElicitationFnT] = None,
124126
credential_key: str | None = None,
125127
):
126128
"""Initializes the McpToolset.
@@ -161,6 +163,11 @@ def __init__(
161163
sampling_callback: Optional callback to handle sampling requests from the
162164
MCP server.
163165
sampling_capabilities: Optional capabilities for sampling.
166+
elicitation_callback: Optional callback to handle elicitation requests
167+
from the MCP server (``elicitation/create``), including URL-mode
168+
elicitations used for out-of-band flows such as auth challenges.
169+
Providing a callback makes the client declare the elicitation
170+
capability during initialization.
164171
credential_key: A user specified key used to load and save this credential
165172
in a credential service. Used with auth_scheme.
166173
"""
@@ -169,6 +176,7 @@ def __init__(
169176

170177
self._sampling_callback = sampling_callback
171178
self._sampling_capabilities = sampling_capabilities
179+
self._elicitation_callback = elicitation_callback
172180

173181
if not connection_params:
174182
raise ValueError("Missing connection params in McpToolset.")
@@ -184,6 +192,7 @@ def __init__(
184192
errlog=self._errlog,
185193
sampling_callback=self._sampling_callback,
186194
sampling_capabilities=self._sampling_capabilities,
195+
elicitation_callback=self._elicitation_callback,
187196
)
188197
self._auth_scheme = auth_scheme
189198
self._auth_credential = auth_credential

src/google/adk/tools/mcp_tool/session_context.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@
2727

2828
from mcp import ClientSession
2929
from mcp import SamplingCapability
30+
from mcp.client.session import ElicitationFnT
3031
from mcp.client.session import SamplingFnT
3132

3233
from ...features import FeatureName
@@ -96,6 +97,7 @@ def __init__(
9697
*,
9798
sampling_callback: Optional[SamplingFnT] = None,
9899
sampling_capabilities: Optional[SamplingCapability] = None,
100+
elicitation_callback: Optional[ElicitationFnT] = None,
99101
):
100102
"""
101103
Args:
@@ -108,6 +110,8 @@ def __init__(
108110
sampling_callback: Optional callback to handle sampling requests from the
109111
MCP server.
110112
sampling_capabilities: Optional capabilities for sampling.
113+
elicitation_callback: Optional callback to handle elicitation requests
114+
from the MCP server (``elicitation/create``).
111115
"""
112116
self._client = client
113117
self._timeout = timeout
@@ -120,6 +124,7 @@ def __init__(
120124
self._task_lock = asyncio.Lock()
121125
self._sampling_callback = sampling_callback
122126
self._sampling_capabilities = sampling_capabilities
127+
self._elicitation_callback = elicitation_callback
123128

124129
@property
125130
def session(self) -> Optional[ClientSession]:
@@ -320,6 +325,7 @@ async def _run(self) -> None:
320325
else None,
321326
sampling_callback=self._sampling_callback,
322327
sampling_capabilities=self._sampling_capabilities,
328+
elicitation_callback=self._elicitation_callback,
323329
)
324330
)
325331
else:
@@ -333,6 +339,7 @@ async def _run(self) -> None:
333339
else None,
334340
sampling_callback=self._sampling_callback,
335341
sampling_capabilities=self._sampling_capabilities,
342+
elicitation_callback=self._elicitation_callback,
336343
)
337344
)
338345
# pylint: disable-next=protected-access

tests/unittests/tools/mcp_tool/test_mcp_session_manager.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -406,6 +406,43 @@ async def test_create_session_stdio_new(self):
406406
# Verify enter_async_context was called (which internally calls __aenter__)
407407
mock_exit_stack.enter_async_context.assert_called_once()
408408

409+
@pytest.mark.asyncio
410+
async def test_create_session_passes_elicitation_callback(self):
411+
"""Elicitation callback is forwarded to the SessionContext."""
412+
413+
async def elicitation_callback(context, params):
414+
return {"action": "decline"}
415+
416+
manager = MCPSessionManager(
417+
self.mock_stdio_connection_params,
418+
elicitation_callback=elicitation_callback,
419+
)
420+
421+
mock_exit_stack = MockAsyncExitStack()
422+
423+
with patch(
424+
"google.adk.tools.mcp_tool.mcp_session_manager.stdio_client"
425+
) as mock_stdio:
426+
with patch(
427+
"google.adk.tools.mcp_tool.mcp_session_manager.AsyncExitStack"
428+
) as mock_exit_stack_class:
429+
with patch(
430+
"google.adk.tools.mcp_tool.mcp_session_manager.SessionContext"
431+
) as mock_session_context_class:
432+
mock_exit_stack_class.return_value = mock_exit_stack
433+
mock_stdio.return_value = AsyncMock()
434+
435+
mock_session = AsyncMock()
436+
mock_session_context = MockSessionContext(session=mock_session)
437+
mock_session_context_class.return_value = mock_session_context
438+
mock_exit_stack.enter_async_context.return_value = mock_session
439+
440+
await manager.create_session()
441+
442+
mock_session_context_class.assert_called_once()
443+
_, kwargs = mock_session_context_class.call_args
444+
assert kwargs["elicitation_callback"] is elicitation_callback
445+
409446
@pytest.mark.asyncio
410447
async def test_create_session_reuse_existing(self):
411448
"""Test reusing an existing connected session."""

tests/unittests/tools/mcp_tool/test_mcp_toolset.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -788,6 +788,34 @@ async def mock_sampling_handler(messages, params=None, context=None):
788788
assert result["role"] == "assistant"
789789
assert result["content"]["text"] == "sampling response"
790790

791+
@pytest.mark.asyncio
792+
async def test_elicitation_callback_plumbed_to_session_manager(self):
793+
"""Elicitation callback reaches the session manager unchanged."""
794+
795+
async def mock_elicitation_handler(context, params):
796+
return {"action": "decline"}
797+
798+
toolset = McpToolset(
799+
connection_params=StreamableHTTPConnectionParams(
800+
url="http://localhost:9999",
801+
timeout=10,
802+
),
803+
elicitation_callback=mock_elicitation_handler,
804+
)
805+
806+
assert toolset._elicitation_callback is mock_elicitation_handler
807+
assert (
808+
toolset._mcp_session_manager._elicitation_callback
809+
is mock_elicitation_handler
810+
)
811+
812+
@pytest.mark.asyncio
813+
async def test_elicitation_callback_defaults_to_none(self):
814+
toolset = McpToolset(connection_params=self.mock_stdio_params)
815+
816+
assert toolset._elicitation_callback is None
817+
assert toolset._mcp_session_manager._elicitation_callback is None
818+
791819
@pytest.mark.asyncio
792820
async def test_get_auth_headers_includes_additional_headers(self):
793821
credential = AuthCredential(

tests/unittests/tools/mcp_tool/test_session_context.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -116,6 +116,36 @@ async def test_start_success_ready_event_set_and_session_returned(self):
116116
# Clean up
117117
await session_context.close()
118118

119+
@pytest.mark.asyncio
120+
async def test_elicitation_callback_passed_to_client_session(self):
121+
"""Elicitation callback is forwarded to the ClientSession."""
122+
123+
async def elicitation_callback(context, params):
124+
return {'action': 'decline'}
125+
126+
mock_client = MockClient()
127+
session_context = SessionContext(
128+
mock_client,
129+
timeout=5.0,
130+
sse_read_timeout=None,
131+
elicitation_callback=elicitation_callback,
132+
)
133+
134+
mock_session = MockClientSession()
135+
136+
with patch(
137+
'google.adk.tools.mcp_tool.session_context.ClientSession'
138+
) as mock_session_class:
139+
mock_session_class.return_value = mock_session
140+
141+
await session_context.start()
142+
143+
mock_session_class.assert_called_once()
144+
_, kwargs = mock_session_class.call_args
145+
assert kwargs['elicitation_callback'] is elicitation_callback
146+
147+
await session_context.close()
148+
119149
@pytest.mark.asyncio
120150
async def test_start_raises_connection_error_on_exception(self):
121151
"""Test that start() raises ConnectionError when exception occurs."""

0 commit comments

Comments
 (0)