feat: Support dynamic per-request headers in MCPToolset

Add a header_provider param which is a callable[ReadonlyContext, Dict[str, Any]] for users to build headers in MCPToolset
fix: https://github.com/google/adk-python/issues/3156
PiperOrigin-RevId: 820412372
This commit is contained in:
Kathy Wu
2025-10-16 15:12:43 -07:00
committed by Copybara-Service
parent 2a8fdd94e1
commit 6dcbb5aca6
8 changed files with 243 additions and 5 deletions
@@ -640,3 +640,74 @@ class TestMCPTool:
with pytest.raises(TypeError):
MCPTool(mcp_tool=self.mock_mcp_tool) # Missing session manager
@pytest.mark.asyncio
async def test_run_async_impl_with_header_provider_no_auth(self):
"""Test running tool with header_provider but no auth."""
expected_headers = {"X-Tenant-ID": "test-tenant"}
header_provider = Mock(return_value=expected_headers)
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
header_provider=header_provider,
)
expected_response = {"result": "success"}
self.mock_session.call_tool = AsyncMock(return_value=expected_response)
tool_context = Mock(spec=ToolContext)
tool_context._invocation_context = Mock()
args = {"param1": "test_value"}
result = await tool._run_async_impl(
args=args, tool_context=tool_context, credential=None
)
assert result == expected_response
header_provider.assert_called_once()
self.mock_session_manager.create_session.assert_called_once_with(
headers=expected_headers
)
self.mock_session.call_tool.assert_called_once_with(
"test_tool", arguments=args
)
@pytest.mark.asyncio
async def test_run_async_impl_with_header_provider_and_oauth2(self):
"""Test running tool with header_provider and OAuth2 auth."""
dynamic_headers = {"X-Tenant-ID": "test-tenant"}
header_provider = Mock(return_value=dynamic_headers)
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
header_provider=header_provider,
)
oauth2_auth = OAuth2Auth(access_token="test_access_token")
credential = AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2, oauth2=oauth2_auth
)
expected_response = {"result": "success"}
self.mock_session.call_tool = AsyncMock(return_value=expected_response)
tool_context = Mock(spec=ToolContext)
tool_context._invocation_context = Mock()
args = {"param1": "test_value"}
result = await tool._run_async_impl(
args=args, tool_context=tool_context, credential=credential
)
assert result == expected_response
header_provider.assert_called_once()
self.mock_session_manager.create_session.assert_called_once()
call_args = self.mock_session_manager.create_session.call_args
headers = call_args[1]["headers"]
assert headers == {
"Authorization": "Bearer test_access_token",
"X-Tenant-ID": "test-tenant",
}
self.mock_session.call_tool.assert_called_once_with(
"test_tool", arguments=args
)
@@ -29,6 +29,7 @@ pytestmark = pytest.mark.skipif(
# Import dependencies with version checking
try:
from google.adk.agents.readonly_context import ReadonlyContext
from google.adk.tools.mcp_tool.mcp_session_manager import MCPSessionManager
from google.adk.tools.mcp_tool.mcp_session_manager import SseConnectionParams
from google.adk.tools.mcp_tool.mcp_session_manager import StdioConnectionParams
@@ -55,6 +56,7 @@ except ImportError as e:
StreamableHTTPConnectionParams = DummyClass
MCPTool = DummyClass
MCPToolset = DummyClass
ReadonlyContext = DummyClass
else:
raise e
@@ -245,6 +247,31 @@ class TestMCPToolset:
assert tools[0].name == "read_file"
assert tools[1].name == "write_file"
@pytest.mark.asyncio
async def test_get_tools_with_header_provider(self):
"""Test get_tools with a header_provider."""
mock_tools = [MockMCPTool("tool1"), MockMCPTool("tool2")]
self.mock_session.list_tools = AsyncMock(
return_value=MockListToolsResult(mock_tools)
)
mock_readonly_context = Mock(spec=ReadonlyContext)
expected_headers = {"X-Tenant-ID": "test-tenant"}
header_provider = Mock(return_value=expected_headers)
toolset = MCPToolset(
connection_params=self.mock_stdio_params,
header_provider=header_provider,
)
toolset._mcp_session_manager = self.mock_session_manager
tools = await toolset.get_tools(readonly_context=mock_readonly_context)
assert len(tools) == 2
header_provider.assert_called_once_with(mock_readonly_context)
self.mock_session_manager.create_session.assert_called_once_with(
headers=expected_headers
)
@pytest.mark.asyncio
async def test_close_success(self):
"""Test successful cleanup."""