feat: Enable MCP Tool Auth (Experimental)

PiperOrigin-RevId: 773002759
This commit is contained in:
Xiang (Sean) Zhou
2025-06-18 11:44:02 -07:00
committed by Copybara-Service
parent 18a541c8fa
commit 157d9be88d
7 changed files with 1229 additions and 101 deletions
@@ -18,9 +18,12 @@ import asyncio
from contextlib import AsyncExitStack
from datetime import timedelta
import functools
import hashlib
import json
import logging
import sys
from typing import Any
from typing import Dict
from typing import Optional
from typing import TextIO
from typing import Union
@@ -105,74 +108,39 @@ class StreamableHTTPConnectionParams(BaseModel):
terminate_on_close: bool = True
def retry_on_closed_resource(session_manager_field_name: str):
"""Decorator to automatically reinitialize session and retry action.
def retry_on_closed_resource(func):
"""Decorator to automatically retry action when MCP session is closed.
When MCP session was closed, the decorator will automatically recreate the
session and retry the action with the same parameters.
Note:
1. session_manager_field_name is the name of the class member field that
contains the MCPSessionManager instance.
2. The session manager must have a reinitialize_session() async method.
Usage:
class MCPTool:
def __init__(self):
self._mcp_session_manager = MCPSessionManager(...)
@retry_on_closed_resource('_mcp_session_manager')
async def use_session(self):
session = await self._mcp_session_manager.create_session()
await session.call_tool()
When MCP session was closed, the decorator will automatically retry the
action once. The create_session method will handle creating a new session
if the old one was disconnected.
Args:
session_manager_field_name: The name of the session manager field.
func: The function to decorate.
Returns:
The decorated function.
"""
def decorator(func):
@functools.wraps(func) # Preserves original function metadata
async def wrapper(self, *args, **kwargs):
try:
return await func(self, *args, **kwargs)
except anyio.ClosedResourceError as close_err:
try:
if hasattr(self, session_manager_field_name):
session_manager = getattr(self, session_manager_field_name)
if hasattr(session_manager, 'reinitialize_session') and callable(
getattr(session_manager, 'reinitialize_session')
):
await session_manager.reinitialize_session()
else:
raise ValueError(
f'Session manager {session_manager_field_name} does not have'
' reinitialize_session method.'
) from close_err
else:
raise ValueError(
f'Session manager field {session_manager_field_name} does not'
' exist in decorated class. Please check the field name in'
' retry_on_closed_resource decorator.'
) from close_err
except Exception as reinit_err:
raise RuntimeError(
f'Error reinitializing: {reinit_err}'
) from reinit_err
return await func(self, *args, **kwargs)
@functools.wraps(func) # Preserves original function metadata
async def wrapper(self, *args, **kwargs):
try:
return await func(self, *args, **kwargs)
except anyio.ClosedResourceError:
# Simply retry the function - create_session will handle
# detecting and replacing disconnected sessions
logger.info('Retrying %s due to closed resource', func.__name__)
return await func(self, *args, **kwargs)
return wrapper
return decorator
return wrapper
class MCPSessionManager:
"""Manages MCP client sessions.
This class provides methods for creating and initializing MCP client sessions,
handling different connection parameters (Stdio and SSE).
handling different connection parameters (Stdio and SSE) and supporting
session pooling based on authentication headers.
"""
def __init__(
@@ -209,30 +177,125 @@ class MCPSessionManager:
else:
self._connection_params = connection_params
self._errlog = errlog
# Each session manager maintains its own exit stack for proper cleanup
self._exit_stack: Optional[AsyncExitStack] = None
self._session: Optional[ClientSession] = None
# Session pool: maps session keys to (session, exit_stack) tuples
self._sessions: Dict[str, tuple[ClientSession, AsyncExitStack]] = {}
# Lock to prevent race conditions in session creation
self._session_lock = asyncio.Lock()
async def create_session(self) -> ClientSession:
def _generate_session_key(
self, merged_headers: Optional[Dict[str, str]] = None
) -> str:
"""Generates a session key based on connection params and merged headers.
For StdioConnectionParams, returns a constant key since headers are not
supported. For SSE and StreamableHTTP connections, generates a key based
on the provided merged headers.
Args:
merged_headers: Already merged headers (base + additional).
Returns:
A unique session key string.
"""
if isinstance(self._connection_params, StdioConnectionParams):
# For stdio connections, headers are not supported, so use constant key
return 'stdio_session'
# For SSE and StreamableHTTP connections, use merged headers
if merged_headers:
headers_json = json.dumps(merged_headers, sort_keys=True)
headers_hash = hashlib.md5(headers_json.encode()).hexdigest()
return f'session_{headers_hash}'
else:
return 'session_no_headers'
def _merge_headers(
self, additional_headers: Optional[Dict[str, str]] = None
) -> Optional[Dict[str, str]]:
"""Merges base connection headers with additional headers.
Args:
additional_headers: Optional headers to merge with connection headers.
Returns:
Merged headers dictionary, or None if no headers are provided.
"""
if isinstance(self._connection_params, StdioConnectionParams) or isinstance(
self._connection_params, StdioServerParameters
):
# Stdio connections don't support headers
return None
base_headers = {}
if (
hasattr(self._connection_params, 'headers')
and self._connection_params.headers
):
base_headers = self._connection_params.headers.copy()
if additional_headers:
base_headers.update(additional_headers)
return base_headers
def _is_session_disconnected(self, session: ClientSession) -> bool:
"""Checks if a session is disconnected or closed.
Args:
session: The ClientSession to check.
Returns:
True if the session is disconnected, False otherwise.
"""
return session._read_stream._closed or session._write_stream._closed
async def create_session(
self, headers: Optional[Dict[str, str]] = None
) -> ClientSession:
"""Creates and initializes an MCP client session.
This method will check if an existing session for the given headers
is still connected. If it's disconnected, it will be cleaned up and
a new session will be created.
Args:
headers: Optional headers to include in the session. These will be
merged with any existing connection headers. Only applicable
for SSE and StreamableHTTP connections.
Returns:
ClientSession: The initialized MCP client session.
"""
# Fast path: if session already exists, return it without acquiring lock
if self._session is not None:
return self._session
# Merge headers once at the beginning
merged_headers = self._merge_headers(headers)
# Generate session key using merged headers
session_key = self._generate_session_key(merged_headers)
# Use async lock to prevent race conditions
async with self._session_lock:
# Double-check: session might have been created while waiting for lock
if self._session is not None:
return self._session
# Check if we have an existing session
if session_key in self._sessions:
session, exit_stack = self._sessions[session_key]
# Create a new exit stack for this session
self._exit_stack = AsyncExitStack()
# Check if the existing session is still connected
if not self._is_session_disconnected(session):
# Session is still good, return it
return session
else:
# Session is disconnected, clean it up
logger.info('Cleaning up disconnected session: %s', session_key)
try:
await exit_stack.aclose()
except Exception as e:
logger.warning('Error during disconnected session cleanup: %s', e)
finally:
del self._sessions[session_key]
# Create a new session (either first time or replacing disconnected one)
exit_stack = AsyncExitStack()
try:
if isinstance(self._connection_params, StdioConnectionParams):
@@ -243,7 +306,7 @@ class MCPSessionManager:
elif isinstance(self._connection_params, SseConnectionParams):
client = sse_client(
url=self._connection_params.url,
headers=self._connection_params.headers,
headers=merged_headers,
timeout=self._connection_params.timeout,
sse_read_timeout=self._connection_params.sse_read_timeout,
)
@@ -252,7 +315,7 @@ class MCPSessionManager:
):
client = streamablehttp_client(
url=self._connection_params.url,
headers=self._connection_params.headers,
headers=merged_headers,
timeout=timedelta(seconds=self._connection_params.timeout),
sse_read_timeout=timedelta(
seconds=self._connection_params.sse_read_timeout
@@ -266,11 +329,11 @@ class MCPSessionManager:
f' {self._connection_params}'
)
transports = await self._exit_stack.enter_async_context(client)
transports = await exit_stack.enter_async_context(client)
# The streamable http client returns a GetSessionCallback in addition to the read/write MemoryObjectStreams
# needed to build the ClientSession, we limit then to the two first values to be compatible with all clients.
if isinstance(self._connection_params, StdioConnectionParams):
session = await self._exit_stack.enter_async_context(
session = await exit_stack.enter_async_context(
ClientSession(
*transports[:2],
read_timeout_seconds=timedelta(
@@ -279,44 +342,38 @@ class MCPSessionManager:
)
)
else:
session = await self._exit_stack.enter_async_context(
session = await exit_stack.enter_async_context(
ClientSession(*transports[:2])
)
await session.initialize()
self._session = session
# Store session and exit stack in the pool
self._sessions[session_key] = (session, exit_stack)
logger.debug('Created new session: %s', session_key)
return session
except Exception:
# If session creation fails, clean up the exit stack
if self._exit_stack:
await self._exit_stack.aclose()
self._exit_stack = None
if exit_stack:
await exit_stack.aclose()
raise
async def close(self):
"""Closes the session and cleans up resources."""
if not self._exit_stack:
return
"""Closes all sessions and cleans up resources."""
async with self._session_lock:
if self._exit_stack:
for session_key in list(self._sessions.keys()):
_, exit_stack = self._sessions[session_key]
try:
await self._exit_stack.aclose()
await exit_stack.aclose()
except Exception as e:
# Log the error but don't re-raise to avoid blocking shutdown
print(
f'Warning: Error during MCP session cleanup: {e}',
'Warning: Error during MCP session cleanup for'
f' {session_key}: {e}',
file=self._errlog,
)
finally:
self._exit_stack = None
self._session = None
async def reinitialize_session(self):
"""Reinitializes the session when connection is lost."""
# Close the old session and create a new one
await self.close()
await self.create_session()
del self._sessions[session_key]
SseServerParams = SseConnectionParams
+76 -12
View File
@@ -14,10 +14,13 @@
from __future__ import annotations
import base64
import json
import logging
from typing import Optional
from google.genai.types import FunctionDeclaration
from google.oauth2.credentials import Credentials
from typing_extensions import override
from .._gemini_schema_util import _to_gemini_schema
@@ -42,13 +45,15 @@ except ImportError as e:
from ...auth.auth_credential import AuthCredential
from ...auth.auth_schemes import AuthScheme
from ..base_tool import BaseTool
from ...auth.auth_tool import AuthConfig
from ..base_authenticated_tool import BaseAuthenticatedTool
# import
from ..tool_context import ToolContext
logger = logging.getLogger("google_adk." + __name__)
class MCPTool(BaseTool):
class MCPTool(BaseAuthenticatedTool):
"""Turns an MCP Tool into an ADK Tool.
Internally, the tool initializes from a MCP Tool, and uses the MCP Session to
@@ -77,19 +82,17 @@ class MCPTool(BaseTool):
Raises:
ValueError: If mcp_tool or mcp_session_manager is None.
"""
if mcp_tool is None:
raise ValueError("mcp_tool cannot be None")
if mcp_session_manager is None:
raise ValueError("mcp_session_manager cannot be None")
super().__init__(
name=mcp_tool.name,
description=mcp_tool.description if mcp_tool.description else "",
auth_config=AuthConfig(
auth_scheme=auth_scheme, raw_auth_credential=auth_credential
)
if auth_scheme
else None,
)
self._mcp_tool = mcp_tool
self._mcp_session_manager = mcp_session_manager
# TODO(cheliu): Support passing auth to MCP Server.
self._auth_scheme = auth_scheme
self._auth_credential = auth_credential
@override
def _get_declaration(self) -> FunctionDeclaration:
@@ -105,8 +108,11 @@ class MCPTool(BaseTool):
)
return function_decl
@retry_on_closed_resource("_mcp_session_manager")
async def run_async(self, *, args, tool_context: ToolContext):
@retry_on_closed_resource
@override
async def _run_async_impl(
self, *, args, tool_context: ToolContext, credential: AuthCredential
):
"""Runs the tool asynchronously.
Args:
@@ -116,8 +122,66 @@ class MCPTool(BaseTool):
Returns:
Any: The response from the tool.
"""
# Extract headers from credential for session pooling
headers = await self._get_headers(tool_context, credential)
# Get the session from the session manager
session = await self._mcp_session_manager.create_session()
session = await self._mcp_session_manager.create_session(headers=headers)
response = await session.call_tool(self.name, arguments=args)
return response
async def _get_headers(
self, tool_context: ToolContext, credential: AuthCredential
) -> Optional[dict[str, str]]:
headers = None
if credential:
if credential.oauth2:
headers = {"Authorization": f"Bearer {credential.oauth2.access_token}"}
elif credential.google_oauth2_json:
google_credential = Credentials.from_authorized_user_info(
json.loads(credential.google_oauth2_json)
)
headers = {"Authorization": f"Bearer {google_credential.token}"}
elif credential.http:
# Handle HTTP authentication schemes
if (
credential.http.scheme.lower() == "bearer"
and credential.http.credentials.token
):
headers = {
"Authorization": f"Bearer {credential.http.credentials.token}"
}
elif credential.http.scheme.lower() == "basic":
# Handle basic auth
if (
credential.http.credentials.username
and credential.http.credentials.password
):
credentials = f"{credential.http.credentials.username}:{credential.http.credentials.password}"
encoded_credentials = base64.b64encode(
credentials.encode()
).decode()
headers = {"Authorization": f"Basic {encoded_credentials}"}
elif credential.http.credentials.token:
# Handle other HTTP schemes with token
headers = {
"Authorization": (
f"{credential.http.scheme} {credential.http.credentials.token}"
)
}
elif credential.api_key:
# For API keys, we'll add them as headers since MCP typically uses header-based auth
# The specific header name would depend on the API, using a common default
# TODO Allow user to specify the header name for API keys.
headers = {"X-API-Key": credential.api_key}
elif credential.service_account:
# Service accounts should be exchanged for access tokens before reaching this point
# If we reach here, we can try to use google_oauth2_json or log a warning
logger.warning(
"Service account credentials should be exchanged for access"
" tokens before MCP session creation"
)
return headers
+11 -1
View File
@@ -22,6 +22,8 @@ from typing import TextIO
from typing import Union
from ...agents.readonly_context import ReadonlyContext
from ...auth.auth_credential import AuthCredential
from ...auth.auth_schemes import AuthScheme
from ..base_tool import BaseTool
from ..base_toolset import BaseToolset
from ..base_toolset import ToolPredicate
@@ -94,6 +96,8 @@ class MCPToolset(BaseToolset):
],
tool_filter: Optional[Union[ToolPredicate, List[str]]] = None,
errlog: TextIO = sys.stderr,
auth_scheme: Optional[AuthScheme] = None,
auth_credential: Optional[AuthCredential] = None,
):
"""Initializes the MCPToolset.
@@ -110,6 +114,8 @@ class MCPToolset(BaseToolset):
list of tool names to include - A ToolPredicate function for custom
filtering logic
errlog: TextIO stream for error logging.
auth_scheme: The auth scheme of the tool for tool calling
auth_credential: The auth credential of the tool for tool calling
"""
super().__init__(tool_filter=tool_filter)
@@ -124,8 +130,10 @@ class MCPToolset(BaseToolset):
connection_params=self._connection_params,
errlog=self._errlog,
)
self._auth_scheme = auth_scheme
self._auth_credential = auth_credential
@retry_on_closed_resource("_mcp_session_manager")
@retry_on_closed_resource
async def get_tools(
self,
readonly_context: Optional[ReadonlyContext] = None,
@@ -151,6 +159,8 @@ class MCPToolset(BaseToolset):
mcp_tool = MCPTool(
mcp_tool=tool,
mcp_session_manager=self._mcp_session_manager,
auth_scheme=self._auth_scheme,
auth_credential=self._auth_credential,
)
if self._is_tool_selected(mcp_tool, readonly_context):
@@ -0,0 +1,13 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
@@ -0,0 +1,342 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import hashlib
from io import StringIO
import json
import sys
from unittest.mock import AsyncMock
from unittest.mock import Mock
from unittest.mock import patch
from google.adk.tools.mcp_tool.mcp_session_manager import MCPSessionManager
from google.adk.tools.mcp_tool.mcp_session_manager import retry_on_closed_resource
from google.adk.tools.mcp_tool.mcp_session_manager import SseConnectionParams
from google.adk.tools.mcp_tool.mcp_session_manager import StdioConnectionParams
from google.adk.tools.mcp_tool.mcp_session_manager import StreamableHTTPConnectionParams
import pytest
# Import real MCP classes
try:
from mcp import StdioServerParameters
except ImportError:
# Create a mock if MCP is not available
class StdioServerParameters:
def __init__(self, command="test_command", args=None):
self.command = command
self.args = args or []
class MockClientSession:
"""Mock ClientSession for testing."""
def __init__(self):
self._read_stream = Mock()
self._write_stream = Mock()
self._read_stream._closed = False
self._write_stream._closed = False
self.initialize = AsyncMock()
class MockAsyncExitStack:
"""Mock AsyncExitStack for testing."""
def __init__(self):
self.aclose = AsyncMock()
self.enter_async_context = AsyncMock()
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
pass
class TestMCPSessionManager:
"""Test suite for MCPSessionManager class."""
def setup_method(self):
"""Set up test fixtures."""
self.mock_stdio_params = StdioServerParameters(
command="test_command", args=[]
)
self.mock_stdio_connection_params = StdioConnectionParams(
server_params=self.mock_stdio_params, timeout=5.0
)
def test_init_with_stdio_server_parameters(self):
"""Test initialization with StdioServerParameters (deprecated)."""
with patch(
"google.adk.tools.mcp_tool.mcp_session_manager.logger"
) as mock_logger:
manager = MCPSessionManager(self.mock_stdio_params)
# Should log deprecation warning
mock_logger.warning.assert_called_once()
assert "StdioServerParameters is not recommended" in str(
mock_logger.warning.call_args
)
# Should convert to StdioConnectionParams
assert isinstance(manager._connection_params, StdioConnectionParams)
assert manager._connection_params.server_params == self.mock_stdio_params
assert manager._connection_params.timeout == 5
def test_init_with_stdio_connection_params(self):
"""Test initialization with StdioConnectionParams."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
assert manager._connection_params == self.mock_stdio_connection_params
assert manager._errlog == sys.stderr
assert manager._sessions == {}
def test_init_with_sse_connection_params(self):
"""Test initialization with SseConnectionParams."""
sse_params = SseConnectionParams(
url="https://example.com/mcp",
headers={"Authorization": "Bearer token"},
timeout=10.0,
)
manager = MCPSessionManager(sse_params)
assert manager._connection_params == sse_params
def test_init_with_streamable_http_params(self):
"""Test initialization with StreamableHTTPConnectionParams."""
http_params = StreamableHTTPConnectionParams(
url="https://example.com/mcp", timeout=15.0
)
manager = MCPSessionManager(http_params)
assert manager._connection_params == http_params
def test_generate_session_key_stdio(self):
"""Test session key generation for stdio connections."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
# For stdio, headers should be ignored and return constant key
key1 = manager._generate_session_key({"Authorization": "Bearer token"})
key2 = manager._generate_session_key(None)
assert key1 == "stdio_session"
assert key2 == "stdio_session"
assert key1 == key2
def test_generate_session_key_sse(self):
"""Test session key generation for SSE connections."""
sse_params = SseConnectionParams(url="https://example.com/mcp")
manager = MCPSessionManager(sse_params)
headers1 = {"Authorization": "Bearer token1"}
headers2 = {"Authorization": "Bearer token2"}
key1 = manager._generate_session_key(headers1)
key2 = manager._generate_session_key(headers2)
key3 = manager._generate_session_key(headers1)
# Different headers should generate different keys
assert key1 != key2
# Same headers should generate same key
assert key1 == key3
# Should be deterministic hash
headers_json = json.dumps(headers1, sort_keys=True)
expected_hash = hashlib.md5(headers_json.encode()).hexdigest()
assert key1 == f"session_{expected_hash}"
def test_merge_headers_stdio(self):
"""Test header merging for stdio connections."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
# Stdio connections don't support headers
headers = manager._merge_headers({"Authorization": "Bearer token"})
assert headers is None
def test_merge_headers_sse(self):
"""Test header merging for SSE connections."""
base_headers = {"Content-Type": "application/json"}
sse_params = SseConnectionParams(
url="https://example.com/mcp", headers=base_headers
)
manager = MCPSessionManager(sse_params)
# With additional headers
additional = {"Authorization": "Bearer token"}
merged = manager._merge_headers(additional)
expected = {
"Content-Type": "application/json",
"Authorization": "Bearer token",
}
assert merged == expected
def test_is_session_disconnected(self):
"""Test session disconnection detection."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
# Create mock session
session = MockClientSession()
# Not disconnected
assert not manager._is_session_disconnected(session)
# Disconnected - read stream closed
session._read_stream._closed = True
assert manager._is_session_disconnected(session)
@pytest.mark.asyncio
async def test_create_session_stdio_new(self):
"""Test creating a new stdio session."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
mock_session = MockClientSession()
mock_exit_stack = MockAsyncExitStack()
with patch(
"google.adk.tools.mcp_tool.mcp_session_manager.stdio_client"
) as mock_stdio:
with patch(
"google.adk.tools.mcp_tool.mcp_session_manager.AsyncExitStack"
) as mock_exit_stack_class:
with patch(
"google.adk.tools.mcp_tool.mcp_session_manager.ClientSession"
) as mock_session_class:
# Setup mocks
mock_exit_stack_class.return_value = mock_exit_stack
mock_stdio.return_value = AsyncMock()
mock_exit_stack.enter_async_context.side_effect = [
("read", "write"), # First call returns transports
mock_session, # Second call returns session
]
mock_session_class.return_value = mock_session
# Create session
session = await manager.create_session()
# Verify session creation
assert session == mock_session
assert len(manager._sessions) == 1
assert "stdio_session" in manager._sessions
# Verify session was initialized
mock_session.initialize.assert_called_once()
@pytest.mark.asyncio
async def test_create_session_reuse_existing(self):
"""Test reusing an existing connected session."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
# Create mock existing session
existing_session = MockClientSession()
existing_exit_stack = MockAsyncExitStack()
manager._sessions["stdio_session"] = (existing_session, existing_exit_stack)
# Session is connected
existing_session._read_stream._closed = False
existing_session._write_stream._closed = False
session = await manager.create_session()
# Should reuse existing session
assert session == existing_session
assert len(manager._sessions) == 1
# Should not create new session
existing_session.initialize.assert_not_called()
@pytest.mark.asyncio
async def test_close_success(self):
"""Test successful cleanup of all sessions."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
# Add mock sessions
session1 = MockClientSession()
exit_stack1 = MockAsyncExitStack()
session2 = MockClientSession()
exit_stack2 = MockAsyncExitStack()
manager._sessions["session1"] = (session1, exit_stack1)
manager._sessions["session2"] = (session2, exit_stack2)
await manager.close()
# All sessions should be closed
exit_stack1.aclose.assert_called_once()
exit_stack2.aclose.assert_called_once()
assert len(manager._sessions) == 0
@pytest.mark.asyncio
async def test_close_with_errors(self):
"""Test cleanup when some sessions fail to close."""
manager = MCPSessionManager(self.mock_stdio_connection_params)
# Add mock sessions
session1 = MockClientSession()
exit_stack1 = MockAsyncExitStack()
exit_stack1.aclose.side_effect = Exception("Close error 1")
session2 = MockClientSession()
exit_stack2 = MockAsyncExitStack()
manager._sessions["session1"] = (session1, exit_stack1)
manager._sessions["session2"] = (session2, exit_stack2)
custom_errlog = StringIO()
manager._errlog = custom_errlog
# Should not raise exception
await manager.close()
# Good session should still be closed
exit_stack2.aclose.assert_called_once()
assert len(manager._sessions) == 0
# Error should be logged
error_output = custom_errlog.getvalue()
assert "Warning: Error during MCP session cleanup" in error_output
assert "Close error 1" in error_output
def test_retry_on_closed_resource_decorator():
"""Test the retry_on_closed_resource decorator."""
call_count = 0
@retry_on_closed_resource
async def mock_function(self):
nonlocal call_count
call_count += 1
if call_count == 1:
import anyio
raise anyio.ClosedResourceError("Resource closed")
return "success"
@pytest.mark.asyncio
async def test_retry():
nonlocal call_count
call_count = 0
mock_self = Mock()
result = await mock_function(mock_self)
assert result == "success"
assert call_count == 2 # First call fails, second succeeds
# Run the test
import asyncio
asyncio.run(test_retry())
@@ -0,0 +1,373 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
from unittest.mock import AsyncMock
from unittest.mock import Mock
from unittest.mock import patch
from google.adk.auth.auth_credential import AuthCredential
from google.adk.auth.auth_credential import AuthCredentialTypes
from google.adk.auth.auth_credential import HttpAuth
from google.adk.auth.auth_credential import HttpCredentials
from google.adk.auth.auth_credential import OAuth2Auth
from google.adk.auth.auth_credential import ServiceAccount
from google.adk.auth.auth_schemes import AuthScheme
from google.adk.auth.auth_schemes import AuthSchemeType
from google.adk.tools.mcp_tool.mcp_session_manager import MCPSessionManager
from google.adk.tools.mcp_tool.mcp_tool import MCPTool
from google.adk.tools.tool_context import ToolContext
from google.genai.types import FunctionDeclaration
import pytest
# Mock MCP Tool from mcp.types
class MockMCPTool:
"""Mock MCP Tool for testing."""
def __init__(self, name="test_tool", description="Test tool description"):
self.name = name
self.description = description
self.inputSchema = {
"type": "object",
"properties": {
"param1": {"type": "string", "description": "First parameter"},
"param2": {"type": "integer", "description": "Second parameter"},
},
"required": ["param1"],
}
class TestMCPTool:
"""Test suite for MCPTool class."""
def setup_method(self):
"""Set up test fixtures."""
self.mock_mcp_tool = MockMCPTool()
self.mock_session_manager = Mock(spec=MCPSessionManager)
self.mock_session = AsyncMock()
self.mock_session_manager.create_session = AsyncMock(
return_value=self.mock_session
)
def test_init_basic(self):
"""Test basic initialization without auth."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
assert tool.name == "test_tool"
assert tool.description == "Test tool description"
assert tool._mcp_tool == self.mock_mcp_tool
assert tool._mcp_session_manager == self.mock_session_manager
def test_init_with_auth(self):
"""Test initialization with authentication."""
# Create real auth scheme instances instead of mocks
from fastapi.openapi.models import OAuth2
auth_scheme = OAuth2(flows={})
auth_credential = AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2,
oauth2=OAuth2Auth(client_id="test_id", client_secret="test_secret"),
)
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
auth_scheme=auth_scheme,
auth_credential=auth_credential,
)
# The auth config is stored in the parent class _credentials_manager
assert tool._credentials_manager is not None
assert tool._credentials_manager._auth_config.auth_scheme == auth_scheme
assert (
tool._credentials_manager._auth_config.raw_auth_credential
== auth_credential
)
def test_init_with_empty_description(self):
"""Test initialization with empty description."""
mock_tool = MockMCPTool(description=None)
tool = MCPTool(
mcp_tool=mock_tool,
mcp_session_manager=self.mock_session_manager,
)
assert tool.description == ""
def test_get_declaration(self):
"""Test function declaration generation."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
declaration = tool._get_declaration()
assert isinstance(declaration, FunctionDeclaration)
assert declaration.name == "test_tool"
assert declaration.description == "Test tool description"
assert declaration.parameters is not None
@pytest.mark.asyncio
async def test_run_async_impl_no_auth(self):
"""Test running tool without authentication."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
# Mock the session response
expected_response = {"result": "success"}
self.mock_session.call_tool = AsyncMock(return_value=expected_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 == expected_response
self.mock_session_manager.create_session.assert_called_once_with(
headers=None
)
# Fix: call_tool uses 'arguments' parameter, not positional args
self.mock_session.call_tool.assert_called_once_with(
"test_tool", arguments=args
)
@pytest.mark.asyncio
async def test_run_async_impl_with_oauth2(self):
"""Test running tool with OAuth2 authentication."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
# Create OAuth2 credential
oauth2_auth = OAuth2Auth(access_token="test_access_token")
credential = AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2, oauth2=oauth2_auth
)
# Mock the session response
expected_response = {"result": "success"}
self.mock_session.call_tool = AsyncMock(return_value=expected_response)
tool_context = Mock(spec=ToolContext)
args = {"param1": "test_value"}
result = await tool._run_async_impl(
args=args, tool_context=tool_context, credential=credential
)
assert result == expected_response
# Check that headers were passed correctly
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"}
@pytest.mark.asyncio
async def test_get_headers_oauth2(self):
"""Test header generation for OAuth2 credentials."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
oauth2_auth = OAuth2Auth(access_token="test_token")
credential = AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2, oauth2=oauth2_auth
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, credential)
assert headers == {"Authorization": "Bearer test_token"}
@pytest.mark.asyncio
async def test_get_headers_http_bearer(self):
"""Test header generation for HTTP Bearer credentials."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
http_auth = HttpAuth(
scheme="bearer", credentials=HttpCredentials(token="bearer_token")
)
credential = AuthCredential(
auth_type=AuthCredentialTypes.HTTP, http=http_auth
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, credential)
assert headers == {"Authorization": "Bearer bearer_token"}
@pytest.mark.asyncio
async def test_get_headers_http_basic(self):
"""Test header generation for HTTP Basic credentials."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
http_auth = HttpAuth(
scheme="basic",
credentials=HttpCredentials(username="user", password="pass"),
)
credential = AuthCredential(
auth_type=AuthCredentialTypes.HTTP, http=http_auth
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, credential)
# Should create Basic auth header with base64 encoded credentials
import base64
expected_encoded = base64.b64encode(b"user:pass").decode()
assert headers == {"Authorization": f"Basic {expected_encoded}"}
@pytest.mark.asyncio
async def test_get_headers_api_key(self):
"""Test header generation for API Key credentials."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
credential = AuthCredential(
auth_type=AuthCredentialTypes.API_KEY, api_key="my_api_key"
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, credential)
assert headers == {"X-API-Key": "my_api_key"}
@pytest.mark.asyncio
@patch("google.adk.tools.mcp_tool.mcp_tool.json")
@patch("google.adk.tools.mcp_tool.mcp_tool.Credentials")
async def test_get_headers_google_oauth2_json(
self, mock_credentials, mock_json
):
"""Test header generation for Google OAuth2 JSON credentials."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
# Mock the JSON parsing and Credentials creation
mock_json.loads.return_value = {"token": "google_token"}
mock_google_credential = Mock()
mock_google_credential.token = "google_access_token"
mock_credentials.from_authorized_user_info.return_value = (
mock_google_credential
)
credential = AuthCredential(
auth_type=AuthCredentialTypes.OAUTH2,
google_oauth2_json='{"token": "google_token"}',
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, credential)
assert headers == {"Authorization": "Bearer google_access_token"}
mock_json.loads.assert_called_once_with('{"token": "google_token"}')
mock_credentials.from_authorized_user_info.assert_called_once_with(
{"token": "google_token"}
)
@pytest.mark.asyncio
async def test_get_headers_no_credential(self):
"""Test header generation with no credentials."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, None)
assert headers is None
@pytest.mark.asyncio
async def test_get_headers_service_account_no_json(self):
"""Test header generation for service account credentials without google_oauth2_json."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
# Create service account credential without google_oauth2_json
service_account = ServiceAccount(scopes=["test"])
credential = AuthCredential(
auth_type=AuthCredentialTypes.SERVICE_ACCOUNT,
service_account=service_account,
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, credential)
# Should return None as no google_oauth2_json is provided
assert headers is None
@pytest.mark.asyncio
async def test_run_async_impl_retry_decorator(self):
"""Test that the retry decorator is applied correctly."""
# This is more of an integration test to ensure the decorator is present
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
# Check that the method has the retry decorator
assert hasattr(tool._run_async_impl, "__wrapped__")
@pytest.mark.asyncio
async def test_get_headers_http_custom_scheme(self):
"""Test header generation for custom HTTP scheme."""
tool = MCPTool(
mcp_tool=self.mock_mcp_tool,
mcp_session_manager=self.mock_session_manager,
)
http_auth = HttpAuth(
scheme="custom", credentials=HttpCredentials(token="custom_token")
)
credential = AuthCredential(
auth_type=AuthCredentialTypes.HTTP, http=http_auth
)
tool_context = Mock(spec=ToolContext)
headers = await tool._get_headers(tool_context, credential)
assert headers == {"Authorization": "custom custom_token"}
def test_init_validation(self):
"""Test that initialization validates required parameters."""
# This test ensures that the MCPTool properly handles its dependencies
with pytest.raises(TypeError):
MCPTool() # Missing required parameters
with pytest.raises(TypeError):
MCPTool(mcp_tool=self.mock_mcp_tool) # Missing session manager
@@ -0,0 +1,269 @@
# Copyright 2025 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from io import StringIO
import sys
from unittest.mock import AsyncMock
from unittest.mock import Mock
from unittest.mock import patch
from google.adk.agents.readonly_context import ReadonlyContext
from google.adk.auth.auth_credential import AuthCredential
from google.adk.auth.auth_schemes import AuthScheme
from google.adk.auth.auth_schemes import AuthSchemeType
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
from google.adk.tools.mcp_tool.mcp_session_manager import StreamableHTTPConnectionParams
from google.adk.tools.mcp_tool.mcp_tool import MCPTool
from google.adk.tools.mcp_tool.mcp_toolset import MCPToolset
import pytest
# Import the real MCP classes for proper instantiation
try:
from mcp import StdioServerParameters
except ImportError:
# Create a mock if MCP is not available
class StdioServerParameters:
def __init__(self, command="test_command", args=None):
self.command = command
self.args = args or []
class MockMCPTool:
"""Mock MCP Tool for testing."""
def __init__(self, name, description="Test tool description"):
self.name = name
self.description = description
self.inputSchema = {
"type": "object",
"properties": {"param": {"type": "string"}},
}
class MockListToolsResult:
"""Mock ListToolsResult for testing."""
def __init__(self, tools):
self.tools = tools
class TestMCPToolset:
"""Test suite for MCPToolset class."""
def setup_method(self):
"""Set up test fixtures."""
self.mock_stdio_params = StdioServerParameters(
command="test_command", args=[]
)
self.mock_session_manager = Mock(spec=MCPSessionManager)
self.mock_session = AsyncMock()
self.mock_session_manager.create_session = AsyncMock(
return_value=self.mock_session
)
def test_init_basic(self):
"""Test basic initialization with StdioServerParameters."""
toolset = MCPToolset(connection_params=self.mock_stdio_params)
# Note: StdioServerParameters gets converted to StdioConnectionParams internally
assert toolset._errlog == sys.stderr
assert toolset._auth_scheme is None
assert toolset._auth_credential is None
def test_init_with_stdio_connection_params(self):
"""Test initialization with StdioConnectionParams."""
stdio_params = StdioConnectionParams(
server_params=self.mock_stdio_params, timeout=10.0
)
toolset = MCPToolset(connection_params=stdio_params)
assert toolset._connection_params == stdio_params
def test_init_with_sse_connection_params(self):
"""Test initialization with SseConnectionParams."""
sse_params = SseConnectionParams(
url="https://example.com/mcp", headers={"Authorization": "Bearer token"}
)
toolset = MCPToolset(connection_params=sse_params)
assert toolset._connection_params == sse_params
def test_init_with_streamable_http_params(self):
"""Test initialization with StreamableHTTPConnectionParams."""
http_params = StreamableHTTPConnectionParams(
url="https://example.com/mcp",
headers={"Content-Type": "application/json"},
)
toolset = MCPToolset(connection_params=http_params)
assert toolset._connection_params == http_params
def test_init_with_tool_filter_list(self):
"""Test initialization with tool filter as list."""
tool_filter = ["tool1", "tool2"]
toolset = MCPToolset(
connection_params=self.mock_stdio_params, tool_filter=tool_filter
)
# The tool filter is stored in the parent BaseToolset class
# We can verify it by checking the filtering behavior in get_tools
assert toolset._is_tool_selected is not None
def test_init_with_auth(self):
"""Test initialization with authentication."""
# Create real auth scheme instances
from fastapi.openapi.models import OAuth2
auth_scheme = OAuth2(flows={})
from google.adk.auth.auth_credential import OAuth2Auth
auth_credential = AuthCredential(
auth_type="oauth2",
oauth2=OAuth2Auth(client_id="test_id", client_secret="test_secret"),
)
toolset = MCPToolset(
connection_params=self.mock_stdio_params,
auth_scheme=auth_scheme,
auth_credential=auth_credential,
)
assert toolset._auth_scheme == auth_scheme
assert toolset._auth_credential == auth_credential
def test_init_missing_connection_params(self):
"""Test initialization with missing connection params raises error."""
with pytest.raises(ValueError, match="Missing connection params"):
MCPToolset(connection_params=None)
@pytest.mark.asyncio
async def test_get_tools_basic(self):
"""Test getting tools without filtering."""
# Mock tools from MCP server
mock_tools = [
MockMCPTool("tool1"),
MockMCPTool("tool2"),
MockMCPTool("tool3"),
]
self.mock_session.list_tools = AsyncMock(
return_value=MockListToolsResult(mock_tools)
)
toolset = MCPToolset(connection_params=self.mock_stdio_params)
toolset._mcp_session_manager = self.mock_session_manager
tools = await toolset.get_tools()
assert len(tools) == 3
for tool in tools:
assert isinstance(tool, MCPTool)
assert tools[0].name == "tool1"
assert tools[1].name == "tool2"
assert tools[2].name == "tool3"
@pytest.mark.asyncio
async def test_get_tools_with_list_filter(self):
"""Test getting tools with list-based filtering."""
# Mock tools from MCP server
mock_tools = [
MockMCPTool("tool1"),
MockMCPTool("tool2"),
MockMCPTool("tool3"),
]
self.mock_session.list_tools = AsyncMock(
return_value=MockListToolsResult(mock_tools)
)
tool_filter = ["tool1", "tool3"]
toolset = MCPToolset(
connection_params=self.mock_stdio_params, tool_filter=tool_filter
)
toolset._mcp_session_manager = self.mock_session_manager
tools = await toolset.get_tools()
assert len(tools) == 2
assert tools[0].name == "tool1"
assert tools[1].name == "tool3"
@pytest.mark.asyncio
async def test_get_tools_with_function_filter(self):
"""Test getting tools with function-based filtering."""
# Mock tools from MCP server
mock_tools = [
MockMCPTool("read_file"),
MockMCPTool("write_file"),
MockMCPTool("list_directory"),
]
self.mock_session.list_tools = AsyncMock(
return_value=MockListToolsResult(mock_tools)
)
def file_tools_filter(tool, context):
"""Filter for file-related tools only."""
return "file" in tool.name
toolset = MCPToolset(
connection_params=self.mock_stdio_params, tool_filter=file_tools_filter
)
toolset._mcp_session_manager = self.mock_session_manager
tools = await toolset.get_tools()
assert len(tools) == 2
assert tools[0].name == "read_file"
assert tools[1].name == "write_file"
@pytest.mark.asyncio
async def test_close_success(self):
"""Test successful cleanup."""
toolset = MCPToolset(connection_params=self.mock_stdio_params)
toolset._mcp_session_manager = self.mock_session_manager
await toolset.close()
self.mock_session_manager.close.assert_called_once()
@pytest.mark.asyncio
async def test_close_with_exception(self):
"""Test cleanup when session manager raises exception."""
toolset = MCPToolset(connection_params=self.mock_stdio_params)
toolset._mcp_session_manager = self.mock_session_manager
# Mock close to raise an exception
self.mock_session_manager.close = AsyncMock(
side_effect=Exception("Cleanup error")
)
custom_errlog = StringIO()
toolset._errlog = custom_errlog
# Should not raise exception
await toolset.close()
# Should log the error
error_output = custom_errlog.getvalue()
assert "Warning: Error during MCPToolset cleanup" in error_output
assert "Cleanup error" in error_output
@pytest.mark.asyncio
async def test_get_tools_retry_decorator(self):
"""Test that get_tools has retry decorator applied."""
toolset = MCPToolset(connection_params=self.mock_stdio_params)
# Check that the method has the retry decorator
assert hasattr(toolset.get_tools, "__wrapped__")