huoyan-enterprise/backend/core/mcp_client.py

108 lines
3.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
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_SERVERSAgent 不会加载 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 客户端已关闭")