diff --git a/src/agents/mcp/util.py b/src/agents/mcp/util.py index 5219f5b9..49b875ab 100644 --- a/src/agents/mcp/util.py +++ b/src/agents/mcp/util.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import copy import functools import inspect @@ -359,10 +360,28 @@ class MCPUtil: try: resolved_meta = await cls._resolve_meta(server, context, tool.name, json_data) merged_meta = cls._merge_mcp_meta(resolved_meta, meta) - if merged_meta is None: - result = await server.call_tool(tool.name, json_data) - else: - result = await server.call_tool(tool.name, json_data, meta=merged_meta) + call_task = asyncio.create_task( + server.call_tool(tool.name, json_data) + if merged_meta is None + else server.call_tool(tool.name, json_data, meta=merged_meta) + ) + try: + done, _ = await asyncio.wait({call_task}, return_when=asyncio.FIRST_COMPLETED) + finished_task = done.pop() + if finished_task.cancelled(): + raise UserError( + f"Failed to call tool '{tool.name}' on MCP server '{server.name}': " + "tool execution was cancelled." + ) + result = finished_task.result() + except asyncio.CancelledError: + if not call_task.done(): + call_task.cancel() + try: + await call_task + except (asyncio.CancelledError, Exception): + pass + raise except UserError: # Re-raise UserError as-is (it already has a good message) raise diff --git a/tests/mcp/test_mcp_util.py b/tests/mcp/test_mcp_util.py index 5a737d07..38db1df2 100644 --- a/tests/mcp/test_mcp_util.py +++ b/tests/mcp/test_mcp_util.py @@ -1,3 +1,4 @@ +import asyncio import dataclasses import json import logging @@ -9,7 +10,7 @@ from mcp.types import CallToolResult, ImageContent, TextContent, Tool as MCPTool from pydantic import BaseModel, TypeAdapter from agents import Agent, FunctionTool, RunContextWrapper, default_tool_error_function -from agents.exceptions import AgentsException, ModelBehaviorError +from agents.exceptions import AgentsException, ModelBehaviorError, UserError from agents.mcp import MCPServer, MCPUtil from agents.tool_context import ToolContext @@ -175,6 +176,46 @@ class CrashingFakeMCPServer(FakeMCPServer): raise Exception("Crash!") +class CancelledFakeMCPServer(FakeMCPServer): + async def call_tool( + self, + tool_name: str, + arguments: dict[str, Any] | None, + meta: dict[str, Any] | None = None, + ): + raise asyncio.CancelledError("synthetic mcp cancel") + + +class SlowFakeMCPServer(FakeMCPServer): + async def call_tool( + self, + tool_name: str, + arguments: dict[str, Any] | None, + meta: dict[str, Any] | None = None, + ): + await asyncio.sleep(60) + return await super().call_tool(tool_name, arguments, meta=meta) + + +class CleanupOnCancelFakeMCPServer(FakeMCPServer): + def __init__(self, cleanup_finished: asyncio.Event): + super().__init__() + self.cleanup_finished = cleanup_finished + + async def call_tool( + self, + tool_name: str, + arguments: dict[str, Any] | None, + meta: dict[str, Any] | None = None, + ): + try: + await asyncio.sleep(60) + except asyncio.CancelledError: + await asyncio.sleep(0.05) + self.cleanup_finished.set() + raise + + @pytest.mark.asyncio async def test_mcp_invocation_crash_causes_error(caplog: pytest.LogCaptureFixture): caplog.set_level(logging.DEBUG) @@ -192,6 +233,159 @@ async def test_mcp_invocation_crash_causes_error(caplog: pytest.LogCaptureFixtur assert "Error invoking MCP tool test_tool_1" in caplog.text +@pytest.mark.asyncio +async def test_mcp_tool_inner_cancellation_becomes_tool_error(): + server = CancelledFakeMCPServer() + server.add_tool("cancel_tool", {}) + + ctx = RunContextWrapper(context=None) + tool = MCPTool(name="cancel_tool", inputSchema={}) + + with pytest.raises(UserError, match="tool execution was cancelled"): + await MCPUtil.invoke_mcp_tool(server, tool, ctx, "{}") + + agent = Agent(name="test-agent") + function_tool = MCPUtil.to_function_tool( + tool, server, convert_schemas_to_strict=False, agent=agent + ) + tool_context = ToolContext( + context=None, + tool_name="cancel_tool", + tool_call_id="test_call_cancelled", + tool_arguments="{}", + ) + + result = await function_tool.on_invoke_tool(tool_context, "{}") + assert isinstance(result, str) + assert "tool execution was cancelled" in result + + +@pytest.mark.asyncio +async def test_mcp_tool_inner_cancellation_still_becomes_tool_error_with_prior_cancel_state(): + current_task = asyncio.current_task() + assert current_task is not None + + current_task.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.sleep(0) + + server = CancelledFakeMCPServer() + server.add_tool("cancel_tool", {}) + + ctx = RunContextWrapper(context=None) + tool = MCPTool(name="cancel_tool", inputSchema={}) + + with pytest.raises(UserError, match="tool execution was cancelled"): + await MCPUtil.invoke_mcp_tool(server, tool, ctx, "{}") + + +@pytest.mark.asyncio +async def test_mcp_tool_outer_cancellation_still_propagates(): + server = SlowFakeMCPServer() + server.add_tool("slow_tool", {}) + + ctx = RunContextWrapper(context=None) + tool = MCPTool(name="slow_tool", inputSchema={}) + + task = asyncio.create_task(MCPUtil.invoke_mcp_tool(server, tool, ctx, "{}")) + await asyncio.sleep(0.05) + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await task + + +@pytest.mark.asyncio +async def test_mcp_tool_outer_cancellation_after_inner_completion_still_propagates( + monkeypatch: pytest.MonkeyPatch, +): + server = FakeMCPServer() + server.add_tool("fast_tool", {}) + + ctx = RunContextWrapper(context=None) + tool = MCPTool(name="fast_tool", inputSchema={}) + + async def fake_wait(tasks, *, return_when): + del return_when + (task,) = tuple(tasks) + await task + raise asyncio.CancelledError("synthetic outer cancellation") + + monkeypatch.setattr(asyncio, "wait", fake_wait) + + with pytest.raises(asyncio.CancelledError): + await MCPUtil.invoke_mcp_tool(server, tool, ctx, "{}") + + +@pytest.mark.asyncio +async def test_mcp_tool_outer_cancellation_after_inner_exception_still_propagates( + monkeypatch: pytest.MonkeyPatch, +): + server = CrashingFakeMCPServer() + server.add_tool("boom_tool", {}) + + ctx = RunContextWrapper(context=None) + tool = MCPTool(name="boom_tool", inputSchema={}) + + async def fake_wait(tasks, *, return_when): + del return_when + (task,) = tuple(tasks) + try: + await task + except Exception: + pass + raise asyncio.CancelledError("synthetic outer cancellation") + + monkeypatch.setattr(asyncio, "wait", fake_wait) + + with pytest.raises(asyncio.CancelledError): + await MCPUtil.invoke_mcp_tool(server, tool, ctx, "{}") + + +@pytest.mark.asyncio +async def test_mcp_tool_outer_cancellation_after_inner_cancellation_still_propagates( + monkeypatch: pytest.MonkeyPatch, +): + server = SlowFakeMCPServer() + server.add_tool("slow_tool", {}) + + ctx = RunContextWrapper(context=None) + tool = MCPTool(name="slow_tool", inputSchema={}) + + async def fake_wait(tasks, *, return_when): + del return_when + (task,) = tuple(tasks) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + raise asyncio.CancelledError("synthetic combined cancellation") + + monkeypatch.setattr(asyncio, "wait", fake_wait) + + with pytest.raises(asyncio.CancelledError): + await MCPUtil.invoke_mcp_tool(server, tool, ctx, "{}") + + +@pytest.mark.asyncio +async def test_mcp_tool_outer_cancellation_waits_for_inner_cleanup(): + cleanup_finished = asyncio.Event() + server = CleanupOnCancelFakeMCPServer(cleanup_finished) + server.add_tool("slow_tool", {}) + + ctx = RunContextWrapper(context=None) + tool = MCPTool(name="slow_tool", inputSchema={}) + + task = asyncio.create_task(MCPUtil.invoke_mcp_tool(server, tool, ctx, "{}")) + await asyncio.sleep(0.05) + task.cancel() + + with pytest.raises(asyncio.CancelledError): + await task + + assert cleanup_finished.is_set() + + @pytest.mark.asyncio async def test_mcp_invocation_mcp_error_reraises(caplog: pytest.LogCaptureFixture): """Test that McpError from server.call_tool is re-raised so the FunctionTool failure