Add support for Chrome MCP (#279)

This commit is contained in:
luijait 2025-09-25 00:10:04 +02:00 committed by GitHub
parent 833c6ed858
commit 82cc20e542
2 changed files with 164 additions and 104 deletions

View File

@ -79,6 +79,9 @@ console = Console()
# Global registry for persistent MCP connections
_GLOBAL_MCP_SERVERS: Dict[str, MCPServer] = {}
# Per-server locks to serialize tool invocations for persistent connections
_SERVER_INVOCATION_LOCKS: Dict[str, asyncio.Lock] = {}
# Global registry for agent-MCP associations
# Maps agent name to list of MCP server names
_AGENT_MCP_ASSOCIATIONS: Dict[str, List[str]] = {}
@ -104,6 +107,7 @@ class GlobalMCPUtil(MCPUtil):
"tool_name": tool.name,
"tool_schema": tool.inputSchema,
"tool_description": tool.description,
"persistent": isinstance(server, MCPServerStdio),
}
# For SSE servers, capture the URL
@ -123,9 +127,9 @@ class GlobalMCPUtil(MCPUtil):
server.params, "encoding_error_handler", "strict"
)
# Create a custom invoke function that creates a new connection each time
# Create a custom invoke function that manages the server lifecycle per invocation
async def invoke_with_fresh_connection(config, context, input_json):
"""Custom invoke function that creates a fresh connection for each invocation"""
"""Invoke an MCP tool, keeping STDIO transports persistent."""
import asyncio
import json
import warnings
@ -133,7 +137,6 @@ class GlobalMCPUtil(MCPUtil):
from cai.sdk.agents.exceptions import AgentsException, ModelBehaviorError
from cai.sdk.agents.mcp import MCPServerSse, MCPServerStdio
# Parse JSON input
try:
json_data = json.loads(input_json) if input_json else {}
except Exception as e:
@ -141,130 +144,167 @@ class GlobalMCPUtil(MCPUtil):
f"Invalid JSON input for tool {config['tool_name']}: {input_json}"
) from e
# Create a fresh server connection with timeout
server = None
result = None
max_retries = 2
retry_count = 0
server = None
should_cleanup = False
persistent = bool(config.get("persistent"))
# Suppress warnings about async generator cleanup
with warnings.catch_warnings():
warnings.filterwarnings("ignore", category=RuntimeWarning)
warnings.filterwarnings("ignore", message=".*asynchronous generator.*")
warnings.filterwarnings("ignore", message=".*ClosedResourceError.*")
try:
if config["type"] == "MCPServerSse":
# Create new SSE server
params = {
"url": config["url"],
"headers": config.get("headers"),
"timeout": config.get("timeout", 5),
"sse_read_timeout": config.get("sse_read_timeout", 60 * 5),
}
# Remove None values
params = {k: v for k, v in params.items() if v is not None}
if persistent:
server_name = config["name"]
server = _GLOBAL_MCP_SERVERS.get(server_name)
if not server or not isinstance(server, MCPServerStdio):
raise AgentsException(
f"MCP server '{server_name}' is unavailable. Use /mcp status to verify it is loaded."
)
server = MCPServerSse(
params,
name=config["name"],
cache_tools_list=False, # Don't cache since it's temporary
lock = _SERVER_INVOCATION_LOCKS.setdefault(
server_name, asyncio.Lock()
)
elif config["type"] == "MCPServerStdio":
# Create new STDIO server
params = {
"command": config["command"],
"args": config.get("args", []),
"env": config.get("env"),
"cwd": config.get("cwd"),
"encoding": config.get("encoding", "utf-8"),
"encoding_error_handler": config.get(
"encoding_error_handler", "strict"
),
}
# Remove None values
params = {k: v for k, v in params.items() if v is not None}
server = MCPServerStdio(params, name=config["name"], cache_tools_list=False)
async with lock:
while retry_count < max_retries:
try:
if not getattr(server, "session", None):
try:
await asyncio.wait_for(server.connect(), timeout=10.0)
except asyncio.TimeoutError:
raise AgentsException(
f"Timeout connecting to MCP server for tool {config['tool_name']}. "
"The server may be down or not responding."
)
result = await asyncio.wait_for(
server.call_tool(config["tool_name"], json_data),
timeout=30.0,
)
break
except asyncio.TimeoutError:
raise AgentsException(
f"Timeout calling MCP tool {config['tool_name']}. "
f"The tool took too long to respond."
)
except Exception:
retry_count += 1
if retry_count >= max_retries:
raise
import logging
logging.debug(
f"Retrying MCP tool {config['tool_name']} (attempt {retry_count}/{max_retries})"
)
try:
await server.cleanup()
except Exception:
pass
server.session = None
await asyncio.sleep(0.5)
else:
raise AgentsException(f"Unknown server type: {config['type']}")
if config["type"] == "MCPServerSse":
params = {
"url": config["url"],
"headers": config.get("headers"),
"timeout": config.get("timeout", 5),
"sse_read_timeout": config.get("sse_read_timeout", 60 * 5),
}
params = {k: v for k, v in params.items() if v is not None}
# Retry logic for connection and tool calls
while retry_count < max_retries:
try:
# Connect to the server with timeout
try:
await asyncio.wait_for(server.connect(), timeout=10.0)
except asyncio.TimeoutError:
raise AgentsException(
f"Timeout connecting to MCP server for tool {config['tool_name']}. "
f"The server may be down or not responding."
)
server = MCPServerSse(
params,
name=config["name"],
cache_tools_list=False,
)
elif config["type"] == "MCPServerStdio":
params = {
"command": config["command"],
"args": config.get("args", []),
"env": config.get("env"),
"cwd": config.get("cwd"),
"encoding": config.get("encoding", "utf-8"),
"encoding_error_handler": config.get(
"encoding_error_handler", "strict"
),
}
params = {k: v for k, v in params.items() if v is not None}
# Call the tool with timeout
server = MCPServerStdio(
params, name=config["name"], cache_tools_list=False
)
else:
raise AgentsException(f"Unknown server type: {config['type']}")
should_cleanup = True
while retry_count < max_retries:
try:
try:
await asyncio.wait_for(server.connect(), timeout=10.0)
except asyncio.TimeoutError:
raise AgentsException(
f"Timeout connecting to MCP server for tool {config['tool_name']}. "
f"The server may be down or not responding."
)
result = await asyncio.wait_for(
server.call_tool(config["tool_name"], json_data), timeout=30.0
server.call_tool(config["tool_name"], json_data),
timeout=30.0,
)
break # Success, exit retry loop
break
except asyncio.TimeoutError:
raise AgentsException(
f"Timeout calling MCP tool {config['tool_name']}. "
f"The tool took too long to respond."
)
except Exception as e:
retry_count += 1
if retry_count >= max_retries:
raise
# Log retry attempt
import logging
logging.debug(f"Retrying MCP tool {config['tool_name']} (attempt {retry_count}/{max_retries})")
# Clear session for SSE servers
if config["type"] == "MCPServerSse" and hasattr(server, 'session'):
server.session = None
await asyncio.sleep(0.5) # Brief delay before retry
except Exception:
retry_count += 1
if retry_count >= max_retries:
raise
import logging
logging.debug(
f"Retrying MCP tool {config['tool_name']} (attempt {retry_count}/{max_retries})"
)
if isinstance(server, MCPServerSse) and hasattr(server, "session"):
server.session = None
await asyncio.sleep(0.5)
except Exception as e:
# Handle ClosedResourceError and connection issues
error_type = type(e).__name__
error_str = str(e).lower()
# Improved error messages for common issues
if (error_type in ("ClosedResourceError", "ExceptionGroup") or
"closedresourceerror" in error_str or
"closed" in error_str or
"connection" in error_str):
if (
error_type in ("ClosedResourceError", "ExceptionGroup")
or "closedresourceerror" in error_str
or "closed" in error_str
or "connection" in error_str
):
raise AgentsException(
f"Connection lost to MCP server for tool {config['tool_name']}. "
f"This is normal for SSE servers. The tool will reconnect automatically "
f"on the next invocation."
"Use /mcp status to reconnect if the issue persists."
) from e
else:
raise AgentsException(
f"Error invoking MCP tool {config['tool_name']}: {type(e).__name__}: {str(e)}"
) from e
raise AgentsException(
f"Error invoking MCP tool {config['tool_name']}: {type(e).__name__}: {str(e)}"
) from e
finally:
# Cleanup the server - handle SSE cleanup issues
if server:
if config["type"] == "MCPServerSse":
# For SSE servers, suppress cleanup errors as they're expected
if should_cleanup and server:
if isinstance(server, MCPServerSse):
try:
# Don't wait too long for SSE cleanup
await asyncio.wait_for(server.cleanup(), timeout=0.5)
except (asyncio.TimeoutError, RuntimeError, Exception):
# Expected for SSE connections - they close abruptly
pass
# Explicitly clear the session to force reconnection next time
server.session = None
else:
# For STDIO servers, cleanup normally
try:
await asyncio.wait_for(server.cleanup(), timeout=5.0)
except (asyncio.TimeoutError, Exception):
pass
# Format the result
if not result:
raise AgentsException(f"No result returned from MCP tool {config['tool_name']}")
@ -336,6 +376,7 @@ def cleanup_mcp_servers():
# Only close the loop if it's not running
if not loop.is_running():
loop.close()
_SERVER_INVOCATION_LOCKS.clear()
except Exception:
pass
@ -987,6 +1028,7 @@ Example: `/mcp add burp 13`
self._run_async(cleanup_server())
del _GLOBAL_MCP_SERVERS[server_name]
_SERVER_INVOCATION_LOCKS.pop(server_name, None)
console.print(f"[green]✓ Removed MCP server '{server_name}'[/green]")
return True
except Exception as e:

View File

@ -69,6 +69,8 @@ class _MCPServerWithClientSession(MCPServer, abc.ABC):
self.session: ClientSession | None = None
self.exit_stack: AsyncExitStack = AsyncExitStack()
self._cleanup_lock: asyncio.Lock = asyncio.Lock()
self._connect_lock: asyncio.Lock = asyncio.Lock()
self._call_lock: asyncio.Lock = asyncio.Lock()
self.cache_tools_list = cache_tools_list
# The cache is always dirty at startup, so that we fetch tools at least once
@ -100,25 +102,32 @@ class _MCPServerWithClientSession(MCPServer, abc.ABC):
async def connect(self):
"""Connect to the server."""
try:
transport = await self.exit_stack.enter_async_context(self.create_streams())
read, write = transport
session = await self.exit_stack.enter_async_context(ClientSession(read, write))
await session.initialize()
self.session = session
except Exception as e:
# Only log connection errors at debug level
error_str = str(e).lower()
error_type = type(e).__name__
if ("connection" in error_str or
"refused" in error_str or
"taskgroup" in error_str or
error_type == "ExceptionGroup"):
logger.debug(f"Expected connection error during MCP server init: {e}")
else:
logger.error(f"Error initializing MCP server: {e}")
await self.cleanup()
raise
if self.session is not None:
return
async with self._connect_lock:
if self.session is not None:
return
try:
transport = await self.exit_stack.enter_async_context(self.create_streams())
read, write = transport
session = await self.exit_stack.enter_async_context(ClientSession(read, write))
await session.initialize()
self.session = session
except Exception as e:
# Only log connection errors at debug level
error_str = str(e).lower()
error_type = type(e).__name__
if ("connection" in error_str or
"refused" in error_str or
"taskgroup" in error_str or
error_type == "ExceptionGroup"):
logger.debug(f"Expected connection error during MCP server init: {e}")
else:
logger.error(f"Error initializing MCP server: {e}")
await self.cleanup()
raise
async def list_tools(self) -> list[MCPTool]:
"""List the tools available on the server."""
@ -141,7 +150,16 @@ class _MCPServerWithClientSession(MCPServer, abc.ABC):
if not self.session:
raise UserError("Server not initialized. Make sure you call `connect()` first.")
return await self.session.call_tool(tool_name, arguments)
async with self._call_lock:
if self.session is None:
await self.connect()
try:
return await self.session.call_tool(tool_name, arguments)
except Exception:
# Ensure resources are released and mark session stale so callers can reconnect
await self.cleanup()
raise
async def cleanup(self):
"""Cleanup the server."""