feat: Add progress_callback support to MCPTool and MCPToolset

Fixes: https://github.com/google/adk-python/issues/3811

Co-authored-by: Xuan Yang <xygoogle@google.com>
PiperOrigin-RevId: 866025995
This commit is contained in:
Xuan Yang
2026-02-05 11:04:36 -08:00
committed by Copybara-Service
parent 9b112e2d13
commit adbc37fea1
7 changed files with 695 additions and 19 deletions
@@ -225,7 +225,7 @@ class TestMCPTool:
)
# Fix: call_tool uses 'arguments' parameter, not positional args
self.mock_session.call_tool.assert_called_once_with(
"test_tool", arguments=args
"test_tool", arguments=args, progress_callback=None
)
@pytest.mark.asyncio
@@ -778,7 +778,7 @@ class TestMCPTool:
headers=expected_headers
)
self.mock_session.call_tool.assert_called_once_with(
"test_tool", arguments=args
"test_tool", arguments=args, progress_callback=None
)
@pytest.mark.asyncio
@@ -821,5 +821,99 @@ class TestMCPTool:
"X-Tenant-ID": "test-tenant",
}
self.mock_session.call_tool.assert_called_once_with(
"test_tool", arguments=args
"test_tool", arguments=args, progress_callback=None
)
def test_init_with_progress_callback(self):
"""Test initialization with progress_callback."""
async def my_progress_callback(
progress: float, total: float | None, message: str | None
) -> None:
pass
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
progress_callback=my_progress_callback,
)
assert tool._progress_callback == my_progress_callback
@pytest.mark.asyncio
async def test_run_async_impl_with_progress_callback(self):
"""Test running tool with progress_callback."""
progress_updates = []
async def my_progress_callback(
progress: float, total: float | None, message: str | None
) -> None:
progress_updates.append((progress, total, message))
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
progress_callback=my_progress_callback,
)
# Mock the session response
mcp_response = CallToolResult(
content=[TextContent(type="text", text="success")]
)
self.mock_session.call_tool = AsyncMock(return_value=mcp_response)
tool_context = Mock(spec=ToolContext)
args = {"param1": "test_value"}
result = await tool._run_async_impl(
args=args, tool_context=tool_context, credential=None
)
assert result == mcp_response.model_dump(exclude_none=True, mode="json")
self.mock_session_manager.create_session.assert_called_once_with(
headers=None
)
# Verify progress_callback was passed to call_tool
self.mock_session.call_tool.assert_called_once_with(
"test_tool", arguments=args, progress_callback=my_progress_callback
)
@pytest.mark.asyncio
async def test_run_async_impl_with_progress_callback_factory(self):
"""Test running tool with progress_callback factory that receives context."""
factory_calls = []
def my_callback_factory(tool_name: str, *, callback_context=None, **kwargs):
factory_calls.append((tool_name, callback_context))
async def callback(
progress: float, total: float | None, message: str | None
) -> None:
pass
return callback
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
progress_callback=my_callback_factory,
)
# Mock the session response
mcp_response = CallToolResult(
content=[TextContent(type="text", text="success")]
)
self.mock_session.call_tool = AsyncMock(return_value=mcp_response)
tool_context = Mock(spec=ToolContext)
args = {"param1": "test_value"}
await tool._run_async_impl(
args=args, tool_context=tool_context, credential=None
)
# Verify factory was called with tool name and tool_context as callback_context
assert len(factory_calls) == 1
assert factory_calls[0][0] == "test_tool"
# callback_context is the tool_context itself (ToolContext extends CallbackContext)
assert factory_calls[0][1] is tool_context
@@ -360,6 +360,101 @@ class TestMcpToolset:
assert tools[0].name == "tool1"
assert tools[1].name == "tool2"
def test_init_with_progress_callback(self):
"""Test initialization with progress_callback."""
async def my_progress_callback(
progress: float, total: float | None, message: str | None
) -> None:
pass
toolset = McpToolset(
connection_params=self.mock_stdio_params,
progress_callback=my_progress_callback,
)
assert toolset._progress_callback == my_progress_callback
@pytest.mark.asyncio
async def test_get_tools_passes_progress_callback_to_mcp_tools(self):
"""Test that get_tools passes progress_callback to created MCPTool instances."""
progress_updates = []
async def my_progress_callback(
progress: float, total: float | None, message: str | None
) -> None:
progress_updates.append((progress, total, message))
mock_tools = [MockMCPTool("tool1"), MockMCPTool("tool2")]
self.mock_session.list_tools = AsyncMock(
return_value=MockListToolsResult(mock_tools)
)
toolset = McpToolset(
connection_params=self.mock_stdio_params,
progress_callback=my_progress_callback,
)
toolset._mcp_session_manager = self.mock_session_manager
tools = await toolset.get_tools()
assert len(tools) == 2
# Verify each tool has the progress_callback set
for tool in tools:
assert tool._progress_callback == my_progress_callback
def test_init_with_progress_callback_factory(self):
"""Test initialization with a ProgressCallbackFactory."""
def my_callback_factory(tool_name: str, *, readonly_context=None, **kwargs):
async def callback(
progress: float, total: float | None, message: str | None
) -> None:
pass
return callback
toolset = McpToolset(
connection_params=self.mock_stdio_params,
progress_callback=my_callback_factory,
)
assert toolset._progress_callback == my_callback_factory
@pytest.mark.asyncio
async def test_get_tools_passes_factory_to_mcp_tools(self):
"""Test that get_tools passes factory directly to MCPTool instances.
The factory is resolved at runtime in McpTool._run_async_impl, not at
tool creation time. This allows the factory to receive ReadonlyContext.
"""
def my_callback_factory(tool_name: str, *, readonly_context=None, **kwargs):
async def callback(
progress: float, total: float | None, message: str | None
) -> None:
pass
return callback
mock_tools = [MockMCPTool("tool1"), MockMCPTool("tool2")]
self.mock_session.list_tools = AsyncMock(
return_value=MockListToolsResult(mock_tools)
)
toolset = McpToolset(
connection_params=self.mock_stdio_params,
progress_callback=my_callback_factory,
)
toolset._mcp_session_manager = self.mock_session_manager
tools = await toolset.get_tools()
assert len(tools) == 2
# Factory is passed directly to each tool (resolved at runtime)
for tool in tools:
assert tool._progress_callback == my_callback_factory
@pytest.mark.asyncio
async def test_list_resources(self):
"""Test listing resources."""