diff --git a/src/google/adk/dependencies/_mcp.py b/src/google/adk/dependencies/_mcp.py index 6e49138016f..aa5f22a102e 100644 --- a/src/google/adk/dependencies/_mcp.py +++ b/src/google/adk/dependencies/_mcp.py @@ -49,6 +49,7 @@ from mcp.client.streamable_http import create_mcp_http_client as create_mcp_http_client from mcp.client.streamable_http import streamable_http_client as streamable_http_client from mcp.server.session import ServerSession as ServerSession +from mcp.types import CallToolResult as CallToolResult from mcp.types import ListResourcesResult as ListResourcesResult from mcp.types import ListToolsResult as ListToolsResult from mcp.types import Tool as Tool @@ -72,6 +73,7 @@ IS_MCP_SDK_V2 = False __all__ = [ + "CallToolResult", "IS_MCP_SDK_V2", "ClientSession", "Context", diff --git a/src/google/adk/tools/mcp_tool/mcp_tool.py b/src/google/adk/tools/mcp_tool/mcp_tool.py index 773637d1d67..57cd31a947b 100644 --- a/src/google/adk/tools/mcp_tool/mcp_tool.py +++ b/src/google/adk/tools/mcp_tool/mcp_tool.py @@ -36,6 +36,7 @@ from ...auth.auth_credential import AuthCredential from ...auth.auth_schemes import AuthScheme from ...auth.auth_tool import AuthConfig +from ...dependencies._mcp import CallToolResult from ...dependencies._mcp import ClientSession from ...dependencies._mcp import IS_MCP_SDK_V2 from ...dependencies._mcp import McpError @@ -672,6 +673,89 @@ async def _run_async_impl( # Resolve progress callback (may be a factory that needs runtime context) resolved_callback = self._resolve_progress_callback(tool_context) + try: + response = await self._call_tool_on_session( + session, final_headers, args, resolved_callback, meta_trace_context + ) + except Exception as e: + if not _is_session_terminated_error(e): + raise + # The server rejected the pooled session id (restart or idle eviction) + # before running the tool, so no side effect happened and one retry on + # a fresh session is safe. `_call_tool_on_session` already discarded + # the dead session from the pool, so `_create_session` builds a new one. + logger.info( + "MCP session was terminated server-side; retrying %s on a fresh" + " session.", + self._mcp_tool.name, + ) + session = await self._create_session(headers=final_headers) + response = await self._call_tool_on_session( + session, final_headers, args, resolved_callback, meta_trace_context + ) + + # Keep the caller's key names off the installed SDK's field naming. + result = _dump_mcp_model(response) + + # 2.x-only field. Acting on it (`input_required` drives elicitation) is a + # feature, not compatibility. Not dropped on 1.x, where a key of that name + # could only be a server extra. + if IS_MCP_SDK_V2: + result.pop("resultType", None) + + # Push UI widget to the event actions if the tool supports it. Dump the + # tool: `payload` is a plain dict, so a model left in it gets serialized by + # whichever sink writes the event, and the sinks disagree -- `inputSchema` + # from those passing `by_alias`, `input_schema` from the session stores. + if self.mcp_app_resource_uri: + # Tests and external subclasses pass duck-typed tools that cannot be + # dumped. Pass those through rather than fail a call that succeeded. + tool_payload: Any = self._mcp_tool + if hasattr(tool_payload, "model_dump"): + tool_payload = _dump_mcp_model(tool_payload) + tool_context.render_ui_widget( + UiWidget( + id=tool_context.function_call_id, + provider="mcp", + payload={ + "resource_uri": self.mcp_app_resource_uri, + "tool": tool_payload, + "tool_args": args, + }, + ) + ) + return result + + async def _call_tool_on_session( + self, + session: ClientSession, + final_headers: dict[str, str] | None, + args: dict[str, Any], + resolved_callback: ProgressFnT | None, + meta_trace_context: dict[str, str] | None, + ) -> CallToolResult: + """Runs one tool call on `session`. + + If the server reports the session as terminated (it restarted or evicted + the session id), the pooled session is discarded before the error + propagates: its local streams and background task still look healthy, so + without this every later call would keep reusing the dead session. + + Args: + session: The pooled session to run the call on. + final_headers: The headers the session was created with, as passed to + ``create_session``. + args: The arguments to pass to the tool. + resolved_callback: The progress callback for this invocation, if any. + meta_trace_context: Trace context to send in the request ``_meta``. + + Returns: + The raw result of the tool call. + + Raises: + Exception: Whatever the call raised; a session-terminated error has + already had its session discarded from the pool. + """ call_coro = session.call_tool( self._mcp_tool.name, arguments=args, @@ -719,41 +803,10 @@ async def _run_async_impl( final_headers, session=session ) raise + return response finally: self._mcp_session_manager._end_session_use(final_headers) # pylint: disable=protected-access - # Keep the caller's key names off the installed SDK's field naming. - result = _dump_mcp_model(response) - - # 2.x-only field. Acting on it (`input_required` drives elicitation) is a - # feature, not compatibility. Not dropped on 1.x, where a key of that name - # could only be a server extra. - if IS_MCP_SDK_V2: - result.pop("resultType", None) - - # Push UI widget to the event actions if the tool supports it. Dump the - # tool: `payload` is a plain dict, so a model left in it gets serialized by - # whichever sink writes the event, and the sinks disagree -- `inputSchema` - # from those passing `by_alias`, `input_schema` from the session stores. - if self.mcp_app_resource_uri: - # Tests and external subclasses pass duck-typed tools that cannot be - # dumped. Pass those through rather than fail a call that succeeded. - tool_payload: Any = self._mcp_tool - if hasattr(tool_payload, "model_dump"): - tool_payload = _dump_mcp_model(tool_payload) - tool_context.render_ui_widget( - UiWidget( - id=tool_context.function_call_id, - provider="mcp", - payload={ - "resource_uri": self.mcp_app_resource_uri, - "tool": tool_payload, - "tool_args": args, - }, - ) - ) - return result - def _detect_error_in_response(self, response: Any) -> str | None: """Telemetry hook: returns an error type if the response indicates an error.""" # `response` is a dumped CallToolResult. `_run_async_impl` restores diff --git a/tests/unittests/tools/mcp_tool/test_mcp_tool.py b/tests/unittests/tools/mcp_tool/test_mcp_tool.py index e67e371d2f7..a2ea0de6af4 100644 --- a/tests/unittests/tools/mcp_tool/test_mcp_tool.py +++ b/tests/unittests/tools/mcp_tool/test_mcp_tool.py @@ -1473,6 +1473,49 @@ async def test_run_async_impl_retries_session_setup(self): assert self.mock_session_manager.create_session.await_count == 2 self.mock_session.call_tool.assert_awaited_once() + @pytest.mark.asyncio + async def test_run_async_impl_recovers_from_terminated_session(self): + """Regression test for https://github.com/google/adk-python/issues/6822. + + A server-side session termination (restart or idle eviction) leaves the + pooled session locally healthy but rejected by the server. The dead + session must be dropped from the pool and the call retried once on a + fresh session: the server refused the request before running the tool, + so the retry cannot duplicate a side effect. + """ + tool = MCPTool( + mcp_tool=self.mock_mcp_tool, + mcp_session_manager=self.mock_session_manager, + ) + dead_session = AsyncMock() + dead_session.call_tool = AsyncMock( + side_effect=make_mcp_error(32600, "Session terminated") + ) + fresh_session = AsyncMock() + fresh_session.call_tool = AsyncMock( + return_value=CallToolResult( + content=[TextContent(type="text", text="ok")] + ) + ) + self.mock_session_manager.create_session = AsyncMock( + side_effect=[dead_session, fresh_session] + ) + tool_context = ToolContext(invocation_context=Mock()) + + result = await tool._run_async_impl( + args={"param1": "test_value"}, + tool_context=tool_context, + credential=None, + ) + + assert result["content"][0]["text"] == "ok" + self.mock_session_manager._discard_session.assert_called_once_with( + None, session=dead_session + ) + assert self.mock_session_manager.create_session.await_count == 2 + dead_session.call_tool.assert_awaited_once() + fresh_session.call_tool.assert_awaited_once() + @pytest.mark.asyncio async def test_get_headers_http_custom_scheme(self): """Test header generation for custom HTTP scheme.""" @@ -2483,9 +2526,15 @@ async def test_run_async_impl_discards_a_session_the_server_dropped(self): args={"param1": "x"}, tool_context=tool_context, credential=None ) - self.mock_session_manager._discard_session.assert_called_once_with( + # The first failure is retried once on a freshly created session; the + # pool hands back the same dead mock, so that attempt fails and is + # discarded too. A third attempt would be a loop. + self.mock_session_manager._discard_session.assert_called_with( None, session=self.mock_session ) + assert self.mock_session_manager._discard_session.call_count == 2 + assert self.mock_session_manager.create_session.await_count == 2 + assert self.mock_session.call_tool.await_count == 2 @pytest.mark.asyncio async def test_run_async_impl_keeps_the_session_when_the_tool_itself_fails(