fix: handle inner MCP tool cancellations as tool errors (#2681)

This commit is contained in:
elainegan-openai
2026-03-15 19:54:17 -07:00
committed by GitHub
parent 8c5c6507ce
commit 710449cb2a
2 changed files with 218 additions and 5 deletions
+23 -4
View File
@@ -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
+195 -1
View File
@@ -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