108 lines
3.0 KiB
Python
108 lines
3.0 KiB
Python
"""
|
||
MCP 客户端管理模块
|
||
|
||
管理 Model Context Protocol 客户端的初始化和获取。
|
||
支持在 .env 中配置多个 MCP 服务地址(见 core.mcp_config)。
|
||
"""
|
||
from typing import Optional
|
||
|
||
from langchain_core.tools import BaseTool
|
||
from langchain_mcp_adapters.client import MultiServerMCPClient
|
||
|
||
from core.config import settings
|
||
from core.mcp_config import build_mcp_connections, parse_mcp_server_entries
|
||
from logger.logging import get_logger
|
||
|
||
logger = get_logger(__name__)
|
||
|
||
# 全局 MCP 客户端
|
||
_mcp_client: Optional[MultiServerMCPClient] = None
|
||
|
||
|
||
def is_mcp_enabled() -> bool:
|
||
"""是否配置了至少一个 MCP 服务地址。"""
|
||
from core.mcp_config import has_mcp_servers
|
||
|
||
return has_mcp_servers(settings)
|
||
|
||
|
||
async def get_mcp_client() -> MultiServerMCPClient:
|
||
"""
|
||
获取或创建全局 MCP 客户端
|
||
|
||
Returns:
|
||
MultiServerMCPClient: MCP 客户端实例
|
||
"""
|
||
global _mcp_client
|
||
|
||
if _mcp_client is None:
|
||
logger.info("初始化 MCP 客户端...")
|
||
|
||
mcp_servers = build_mcp_connections(settings)
|
||
entries = parse_mcp_server_entries(settings)
|
||
|
||
if not mcp_servers:
|
||
logger.warning(
|
||
"未配置 MCP 服务(MCP_URL / MCP_URLS / MCP_SERVERS),Agent 不会加载 MCP 工具"
|
||
)
|
||
|
||
# 多服务器时为工具名加前缀,避免 excel_get_schema 等重名冲突
|
||
_mcp_client = MultiServerMCPClient(
|
||
mcp_servers,
|
||
tool_name_prefix=len(entries) > 1,
|
||
)
|
||
logger.info(
|
||
"MCP 客户端初始化完成,共 {} 个服务器",
|
||
len(mcp_servers),
|
||
)
|
||
|
||
return _mcp_client
|
||
|
||
|
||
async def load_mcp_tools_safe() -> list[BaseTool]:
|
||
"""
|
||
加载 MCP 工具;单个服务器连接/鉴权失败时跳过并继续,不阻断聊天。
|
||
"""
|
||
if not is_mcp_enabled():
|
||
return []
|
||
|
||
client = await get_mcp_client()
|
||
if not client.connections:
|
||
return []
|
||
|
||
all_tools: list[BaseTool] = []
|
||
for server_name in client.connections:
|
||
conn = client.connections[server_name]
|
||
url = conn.get("url", "?")
|
||
try:
|
||
tools = await client.get_tools(server_name=server_name)
|
||
all_tools.extend(tools)
|
||
logger.info(
|
||
"MCP 服务器 {} ({}) 已加载 {} 个工具",
|
||
server_name,
|
||
url,
|
||
len(tools),
|
||
)
|
||
except Exception as exc:
|
||
logger.warning(
|
||
"MCP 服务器 {} ({}) 不可用,已跳过: {}",
|
||
server_name,
|
||
url,
|
||
exc,
|
||
)
|
||
|
||
if not all_tools and client.connections:
|
||
logger.warning("所有 MCP 服务器均不可用,聊天将继续但不提供 MCP 工具")
|
||
|
||
return all_tools
|
||
|
||
|
||
async def close_mcp_client():
|
||
"""关闭 MCP 客户端"""
|
||
global _mcp_client
|
||
|
||
if _mcp_client is not None:
|
||
logger.info("关闭 MCP 客户端...")
|
||
_mcp_client = None
|
||
logger.info("MCP 客户端已关闭")
|