Files
xiaozhi-esp32-server/main/xiaozhi-server/core/mcp/MCPClient.py
T

144 lines
5.0 KiB
Python
Raw Normal View History

from __future__ import annotations
2025-03-20 18:20:09 +08:00
from datetime import timedelta
import asyncio, os, shutil, concurrent.futures
2025-03-20 08:59:45 +08:00
from contextlib import AsyncExitStack
from typing import Optional, List, Dict, Any
2025-03-20 08:59:45 +08:00
from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client
2025-05-04 23:33:07 +08:00
from mcp.client.sse import sse_client
2025-03-20 08:59:45 +08:00
from config.logger import setup_logging
2025-06-07 13:46:35 +08:00
from core.utils.util import sanitize_tool_name
2025-03-20 08:59:45 +08:00
TAG = __name__
2025-03-20 08:59:45 +08:00
class MCPClient:
def __init__(self, config: Dict[str, Any]):
2025-03-20 08:59:45 +08:00
self.logger = setup_logging()
self.config = config
self._worker_task: Optional[asyncio.Task] = None
self._ready_evt = asyncio.Event()
self._shutdown_evt = asyncio.Event()
self.session: Optional[ClientSession] = None
2025-06-07 13:46:35 +08:00
self.tools: List = [] # original tool objects
self.tools_dict: Dict[str, Any] = {}
self.name_mapping: Dict[str, str] = {}
2025-03-20 08:59:45 +08:00
async def initialize(self):
if self._worker_task:
return
self._worker_task = asyncio.create_task(self._worker(), name="MCPClientWorker")
await self._ready_evt.wait()
2025-03-20 08:59:45 +08:00
self.logger.bind(tag=TAG).info(
2025-06-07 13:46:35 +08:00
f"Connected, tools = {[name for name in self.name_mapping.values()]}"
)
2025-03-20 08:59:45 +08:00
async def cleanup(self):
if not self._worker_task:
return
self._shutdown_evt.set()
try:
await asyncio.wait_for(self._worker_task, timeout=20)
except (asyncio.TimeoutError, Exception) as e:
self.logger.bind(tag=TAG).error(f"worker shutdown err: {e}")
finally:
self._worker_task = None
def has_tool(self, name: str) -> bool:
2025-06-07 13:46:35 +08:00
return name in self.tools_dict
def get_available_tools(self):
return [
{
"type": "function",
"function": {
2025-06-07 13:46:35 +08:00
"name": name,
"description": tool.description,
"parameters": tool.inputSchema,
},
}
2025-06-07 13:46:35 +08:00
for name, tool in self.tools_dict.items()
]
async def call_tool(self, name: str, args: dict):
if not self.session:
raise RuntimeError("MCPClient not initialized")
2025-06-07 13:46:35 +08:00
real_name = self.name_mapping.get(name, name)
loop = self._worker_task.get_loop()
2025-06-07 13:46:35 +08:00
coro = self.session.call_tool(real_name, args)
if loop is asyncio.get_running_loop():
return await coro
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
return await asyncio.wrap_future(fut)
async def _worker(self):
async with AsyncExitStack() as stack:
try:
# 建立 StdioClient
if "command" in self.config:
cmd = (
shutil.which("npx")
if self.config["command"] == "npx"
else self.config["command"]
)
env = {**os.environ, **self.config.get("env", {})}
params = StdioServerParameters(
command=cmd,
args=self.config.get("args", []),
env=env,
)
2025-05-31 10:14:23 +08:00
stdio_r, stdio_w = await stack.enter_async_context(
stdio_client(params)
)
read_stream, write_stream = stdio_r, stdio_w
# 建立SSEClient
elif "url" in self.config:
2025-05-30 11:08:32 +00:00
if "API_ACCESS_TOKEN" in self.config:
2025-05-31 10:14:23 +08:00
headers = {
"Authorization": f"Bearer {self.config['API_ACCESS_TOKEN']}"
}
2025-05-30 11:08:32 +00:00
else:
headers = {}
2025-05-31 10:14:23 +08:00
sse_r, sse_w = await stack.enter_async_context(
sse_client(self.config["url"], headers=headers)
)
read_stream, write_stream = sse_r, sse_w
else:
raise ValueError("MCPClient config must include 'command' or 'url'")
self.session = await stack.enter_async_context(
ClientSession(
read_stream=read_stream,
write_stream=write_stream,
read_timeout_seconds=timedelta(seconds=15),
)
)
await self.session.initialize()
# 获取工具
self.tools = (await self.session.list_tools()).tools
2025-06-07 13:46:35 +08:00
for t in self.tools:
sanitized = sanitize_tool_name(t.name)
self.tools_dict[sanitized] = t
self.name_mapping[sanitized] = t.name
self._ready_evt.set()
# 挂起等待关闭
await self._shutdown_evt.wait()
except Exception as e:
self.logger.bind(tag=TAG).error(f"worker error: {e}")
self._ready_evt.set()
raise