mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-28 10:03:54 +08:00
resolve: merge upstream/main into feature/mcp-exit-fix and fix conflicts
This commit is contained in:
@@ -1,67 +1,52 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio, os, shutil, concurrent.futures
|
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
|
import asyncio, os, shutil, concurrent.futures
|
||||||
from contextlib import AsyncExitStack
|
from contextlib import AsyncExitStack
|
||||||
from typing import Optional, List
|
from typing import Optional, List, Dict, Any
|
||||||
|
|
||||||
from mcp import ClientSession, StdioServerParameters
|
from mcp import ClientSession, StdioServerParameters
|
||||||
from mcp.client.stdio import stdio_client
|
from mcp.client.stdio import stdio_client
|
||||||
|
from mcp.client.sse import sse_client
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
|
|
||||||
class MCPClient:
|
class MCPClient:
|
||||||
def __init__(self, config: dict):
|
def __init__(self, config: Dict[str, Any]):
|
||||||
self.logger = setup_logging()
|
self.logger = setup_logging()
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
# Back‑worker task & 状态同步
|
|
||||||
self._worker_task: Optional[asyncio.Task] = None
|
self._worker_task: Optional[asyncio.Task] = None
|
||||||
self._ready_evt = asyncio.Event()
|
self._ready_evt = asyncio.Event()
|
||||||
self._shutdown_evt = asyncio.Event()
|
self._shutdown_evt = asyncio.Event()
|
||||||
|
|
||||||
# 运行时资源
|
|
||||||
self.session: Optional[ClientSession] = None
|
self.session: Optional[ClientSession] = None
|
||||||
self.tools: List = []
|
self.tools: List = []
|
||||||
|
|
||||||
async def initialize(self):
|
async def initialize(self):
|
||||||
"""
|
|
||||||
启动后台 task,并等待其就绪(拿到 `tools`)。
|
|
||||||
"""
|
|
||||||
if self._worker_task:
|
if self._worker_task:
|
||||||
return # 已经 init 过
|
return
|
||||||
|
|
||||||
# 在当前 loop 创建后台 task
|
|
||||||
self._worker_task = asyncio.create_task(self._worker(), name="MCPClientWorker")
|
self._worker_task = asyncio.create_task(self._worker(), name="MCPClientWorker")
|
||||||
await self._ready_evt.wait() # 等待 worker 初始化完成
|
await self._ready_evt.wait()
|
||||||
|
|
||||||
# 此时 tools 已填充
|
|
||||||
self.logger.bind(tag=TAG).info(
|
self.logger.bind(tag=TAG).info(
|
||||||
f"Connected, tools = {[t.name for t in self.tools]}"
|
f"Connected, tools = {[t.name for t in self.tools]}"
|
||||||
)
|
)
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
"""
|
|
||||||
对外关闭接口:
|
|
||||||
· 只负责发出 “关机信号”
|
|
||||||
· 等待后台 task 正常退出
|
|
||||||
在任何 loop / task 调用都安全。
|
|
||||||
"""
|
|
||||||
if not self._worker_task:
|
if not self._worker_task:
|
||||||
return
|
return
|
||||||
|
|
||||||
self._shutdown_evt.set() # 发信号
|
self._shutdown_evt.set()
|
||||||
try:
|
try:
|
||||||
await asyncio.wait_for(self._worker_task, timeout=15)
|
await asyncio.wait_for(self._worker_task, timeout=20)
|
||||||
except (asyncio.TimeoutError, Exception) as e:
|
except (asyncio.TimeoutError, Exception) as e:
|
||||||
self.logger.bind(tag=TAG).error(f"worker shutdown err: {e}")
|
self.logger.bind(tag=TAG).error(f"worker shutdown err: {e}")
|
||||||
finally:
|
finally:
|
||||||
self._worker_task = None
|
self._worker_task = None
|
||||||
|
|
||||||
# ----------------------------- 工具接口 -----------------------------
|
|
||||||
|
|
||||||
def has_tool(self, name: str) -> bool:
|
def has_tool(self, name: str) -> bool:
|
||||||
return any(t.name == name for t in self.tools)
|
return any(t.name == name for t in self.tools)
|
||||||
|
|
||||||
@@ -79,67 +64,62 @@ class MCPClient:
|
|||||||
]
|
]
|
||||||
|
|
||||||
async def call_tool(self, name: str, args: dict):
|
async def call_tool(self, name: str, args: dict):
|
||||||
"""
|
if not self.session:
|
||||||
转发到 session.call_tool。
|
|
||||||
若在 worker 之外的 task 调用,会通过 run_coroutine_threadsafe
|
|
||||||
投递到 worker 所在 loop 中执行,保证线程安全。
|
|
||||||
"""
|
|
||||||
if not self.session: # 尚未就绪
|
|
||||||
raise RuntimeError("MCPClient not initialized")
|
raise RuntimeError("MCPClient not initialized")
|
||||||
|
|
||||||
loop = self._worker_task.get_loop()
|
loop = self._worker_task.get_loop()
|
||||||
coro = self.session.call_tool(name, args)
|
coro = self.session.call_tool(name, args)
|
||||||
|
|
||||||
# 在同一个 loop ➜ 直接 await
|
|
||||||
if loop is asyncio.get_running_loop():
|
if loop is asyncio.get_running_loop():
|
||||||
return await coro
|
return await coro
|
||||||
|
|
||||||
# 跨 loop ➜ run_coroutine_threadsafe
|
|
||||||
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
|
fut: concurrent.futures.Future = asyncio.run_coroutine_threadsafe(coro, loop)
|
||||||
return await asyncio.wrap_future(fut)
|
return await asyncio.wrap_future(fut)
|
||||||
|
|
||||||
# ----------------------------- 后台 task -----------------------------
|
|
||||||
|
|
||||||
async def _worker(self):
|
async def _worker(self):
|
||||||
"""
|
|
||||||
单线程协程:
|
|
||||||
1. 创建所有异步资源
|
|
||||||
2. set_ready → 供外部使用
|
|
||||||
3. 等待 shutdown_evt
|
|
||||||
4. 自动随 AsyncExitStack 退出而清理资源
|
|
||||||
"""
|
|
||||||
async with AsyncExitStack() as stack:
|
async with AsyncExitStack() as stack:
|
||||||
try:
|
try:
|
||||||
# ---------- 启动后端进程 ----------
|
# 建立 StdioClient
|
||||||
cmd = shutil.which("npx") if self.config["command"] == "npx" else self.config["command"]
|
if "command" in self.config:
|
||||||
env = {**os.environ, **self.config.get("env", {})}
|
cmd = (
|
||||||
params = StdioServerParameters(
|
shutil.which("npx")
|
||||||
command=cmd,
|
if self.config["command"] == "npx"
|
||||||
args=self.config.get("args", []),
|
else self.config["command"]
|
||||||
env=env,
|
)
|
||||||
)
|
env = {**os.environ, **self.config.get("env", {})}
|
||||||
stdio_r, stdio_w = await stack.enter_async_context(stdio_client(params))
|
params = StdioServerParameters(
|
||||||
|
command=cmd,
|
||||||
|
args=self.config.get("args", []),
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
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:
|
||||||
|
sse_r, sse_w = await stack.enter_async_context(sse_client(self.config["url"]))
|
||||||
|
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(
|
self.session = await stack.enter_async_context(
|
||||||
ClientSession(
|
ClientSession(
|
||||||
read_stream=stdio_r,
|
read_stream=read_stream,
|
||||||
write_stream=stdio_w,
|
write_stream=write_stream,
|
||||||
read_timeout_seconds=timedelta(seconds=15),
|
read_timeout_seconds=timedelta(seconds=15),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
await self.session.initialize()
|
await self.session.initialize()
|
||||||
|
|
||||||
# ---------- 工具 ----------
|
# 获取工具
|
||||||
self.tools = (await self.session.list_tools()).tools
|
self.tools = (await self.session.list_tools()).tools
|
||||||
|
|
||||||
# 初始化完成,放行外部
|
|
||||||
self._ready_evt.set()
|
self._ready_evt.set()
|
||||||
|
|
||||||
# ---------- 挂起等待关闭 ----------
|
# 挂起等待关闭
|
||||||
await self._shutdown_evt.wait()
|
await self._shutdown_evt.wait()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.bind(tag=TAG).error(f"worker error: {e}")
|
self.logger.bind(tag=TAG).error(f"worker error: {e}")
|
||||||
self._ready_evt.set() # 确保外部不会卡死
|
self._ready_evt.set()
|
||||||
raise
|
raise
|
||||||
|
|||||||
@@ -50,9 +50,9 @@ class MCPManager:
|
|||||||
"""初始化所有MCP服务"""
|
"""初始化所有MCP服务"""
|
||||||
config = self.load_config()
|
config = self.load_config()
|
||||||
for name, srv_config in config.items():
|
for name, srv_config in config.items():
|
||||||
if not srv_config.get("command"):
|
if not srv_config.get("command") and not srv_config.get("url"):
|
||||||
self.logger.bind(tag=TAG).warning(
|
self.logger.bind(tag=TAG).warning(
|
||||||
f"Skipping server {name}: command not specified"
|
f"Skipping server {name}: neither command nor url specified"
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
|||||||
@@ -3,7 +3,8 @@
|
|||||||
"在data目录下创建.mcp_server_settings.json文件,可以选择下面的MCP服务,也可以自行添加新的MCP服务。",
|
"在data目录下创建.mcp_server_settings.json文件,可以选择下面的MCP服务,也可以自行添加新的MCP服务。",
|
||||||
"后面不断测试补充好用的mcp服务,欢迎大家一起补充。",
|
"后面不断测试补充好用的mcp服务,欢迎大家一起补充。",
|
||||||
"记得删除注释行,des属性仅为说明,不会被解析。",
|
"记得删除注释行,des属性仅为说明,不会被解析。",
|
||||||
"des和link属性,仅为说明安装方式,方便大家查看原始链接,不是必须项。"
|
"des和link属性,仅为说明安装方式,方便大家查看原始链接,不是必须项。",
|
||||||
|
"当前支持stdio/sse两种模式。"
|
||||||
],
|
],
|
||||||
"mcpServers": {
|
"mcpServers": {
|
||||||
"filesystem": {
|
"filesystem": {
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ mem0ai==0.1.62
|
|||||||
bs4==0.0.2
|
bs4==0.0.2
|
||||||
modelscope==1.23.2
|
modelscope==1.23.2
|
||||||
sherpa_onnx==1.11.0
|
sherpa_onnx==1.11.0
|
||||||
mcp==1.4.1
|
mcp==1.7.1
|
||||||
cnlunar==0.2.0
|
cnlunar==0.2.0
|
||||||
PySocks==1.7.1
|
PySocks==1.7.1
|
||||||
dashscope==1.23.1
|
dashscope==1.23.1
|
||||||
|
|||||||
Reference in New Issue
Block a user