mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: Enable MCP Tool Auth (Experimental)
PiperOrigin-RevId: 773002759
This commit is contained in:
committed by
Copybara-Service
parent
18a541c8fa
commit
157d9be88d
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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__")
|
||||
Reference in New Issue
Block a user