mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
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:
committed by
Copybara-Service
parent
2a8fdd94e1
commit
6dcbb5aca6
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user