refactor: introduce transport-neutral connection runtime

This commit is contained in:
caixypromise
2026-07-27 02:05:26 +08:00
parent f5ed1aaec8
commit 0c582ed3b6
56 changed files with 9931 additions and 237 deletions
@@ -127,7 +127,16 @@ class DeviceIoTExecutor(ToolExecutor):
send_message = json.dumps(
{"type": "iot", "commands": [command]}
)
await self.conn.websocket.send(send_message)
# 使用transport接口发送消息
if hasattr(self.conn, 'transport') and self.conn.transport:
await self.conn.transport.send(send_message)
elif hasattr(self.conn, 'websocket') and self.conn.websocket:
# 兼容旧版本
logger.warning("未找到SessionContext的传输层接口, 回退使用旧版conn.websocket发送消息")
await self.conn.websocket.send(send_message)
else:
raise AttributeError("无法找到可用的传输层接口")
return
raise Exception(f"未找到设备{device_name}的方法{method_name}")
@@ -17,7 +17,7 @@ class MCPClient:
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.next_id = 10000
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
@@ -91,3 +91,12 @@ class MCPClient:
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
async def close(self):
async with self.lock:
pending = list(self.call_results.values())
self.call_results.clear()
self.ready = False
for future in pending:
if not future.done():
future.set_exception(ConnectionError("MCP会话已关闭"))
@@ -3,11 +3,11 @@
import json
import asyncio
import re
from concurrent.futures import Future
from core.utils.util import get_vision_url, sanitize_tool_name
from core.utils.util import get_vision_url
from core.utils.auth import AuthToken
from config.logger import setup_logging
from typing import TYPE_CHECKING
from .mcp_client import MCPClient
if TYPE_CHECKING:
from core.connection import ConnectionHandler
@@ -15,108 +15,45 @@ if TYPE_CHECKING:
TAG = __name__
logger = setup_logging()
class MCPClient:
"""设备端MCP客户端,用于管理MCP状态和工具"""
def __init__(self):
self.tools = {} # sanitized_name -> tool_data
self.name_mapping = {}
self.ready = False
self.call_results = {} # To store Futures for tool call responses
self.next_id = 1
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
def has_tool(self, name: str) -> bool:
return name in self.tools
def get_available_tools(self) -> list:
# Check if the cache is valid
if self._cached_available_tools is not None:
return self._cached_available_tools
# If cache is not valid, regenerate the list
result = []
for tool_name, tool_data in self.tools.items():
function_def = {
"name": tool_name,
"description": tool_data["description"],
"parameters": {
"type": tool_data["inputSchema"].get("type", "object"),
"properties": tool_data["inputSchema"].get("properties", {}),
"required": tool_data["inputSchema"].get("required", []),
},
}
result.append({"type": "function", "function": function_def})
self._cached_available_tools = result # Store the generated list in cache
return result
async def is_ready(self) -> bool:
async with self.lock:
return self.ready
async def set_ready(self, status: bool):
async with self.lock:
self.ready = status
async def add_tool(self, tool_data: dict):
async with self.lock:
sanitized_name = sanitize_tool_name(tool_data["name"])
self.tools[sanitized_name] = tool_data
self.name_mapping[sanitized_name] = tool_data["name"]
self._cached_available_tools = (
None # Invalidate the cache when a tool is added
)
async def get_next_id(self) -> int:
async with self.lock:
current_id = self.next_id
self.next_id += 1
return current_id
async def register_call_result_future(self, id: int, future: Future):
async with self.lock:
self.call_results[id] = future
async def resolve_call_result(self, id: int, result: any):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_result(result)
async def reject_call_result(self, id: int, exception: Exception):
async with self.lock:
if id in self.call_results:
future = self.call_results.pop(id)
if not future.done():
future.set_exception(exception)
async def cleanup_call_result(self, id: int):
async with self.lock:
if id in self.call_results:
self.call_results.pop(id)
async def send_mcp_message(conn: "ConnectionHandler", payload: dict):
async def send_mcp_message(
conn: "ConnectionHandler",
payload: dict,
transport=None,
*,
raise_on_error: bool = False,
):
"""Helper to send MCP messages, encapsulating common logic."""
if not conn.features.get("mcp"):
features = getattr(conn, "features", {}) or {}
if not features.get("mcp"):
logger.bind(tag=TAG).warning("客户端不支持MCP,无法发送MCP消息")
return
message = json.dumps({"type": "mcp", "payload": payload})
try:
await conn.websocket.send(message)
# 优先使用传入的transport,否则尝试从conn获取
if transport:
await transport.send(message)
elif getattr(conn, "transport", None):
# 新架构
await conn.transport.send(message)
elif getattr(conn, "websocket", None):
# 兼容旧版本
await conn.websocket.send(message)
else:
raise AttributeError("无法找到可用的传输层接口")
logger.bind(tag=TAG).debug(f"成功发送MCP消息: {message}")
except Exception as e:
logger.bind(tag=TAG).error(f"发送MCP消息失败: {e}")
if raise_on_error:
raise
async def handle_mcp_message(
conn: "ConnectionHandler", mcp_client: MCPClient, payload: dict
conn: "ConnectionHandler",
mcp_client: MCPClient,
payload: dict,
transport=None,
):
"""处理MCP消息,包括初始化、工具列表和工具调用响应等"""
logger.bind(tag=TAG).debug(f"处理MCP消息: {str(payload)[:100]}")
@@ -150,7 +87,7 @@ async def handle_mcp_message(
await asyncio.sleep(1)
logger.bind(tag=TAG).debug("初始化完成,开始请求MCP工具列表")
await send_mcp_tools_list_request(conn)
await send_mcp_tools_list_request(conn, transport)
return
@@ -207,7 +144,7 @@ async def handle_mcp_message(
next_cursor = result.get("nextCursor", "")
if next_cursor:
logger.bind(tag=TAG).debug(f"有更多工具,nextCursor: {next_cursor}")
await send_mcp_tools_list_continue_request(conn, next_cursor)
await send_mcp_tools_list_continue_request(conn, next_cursor, transport)
else:
await mcp_client.set_ready(True)
logger.bind(tag=TAG).debug("所有工具已获取,MCP客户端准备就绪")
@@ -235,7 +172,9 @@ async def handle_mcp_message(
)
async def send_mcp_initialize_message(conn: "ConnectionHandler"):
async def send_mcp_initialize_message(
conn: "ConnectionHandler", transport=None
):
"""发送MCP初始化消息"""
vision_url = get_vision_url(conn.config)
@@ -267,10 +206,12 @@ async def send_mcp_initialize_message(conn: "ConnectionHandler"):
},
}
logger.bind(tag=TAG).debug("发送MCP初始化消息")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def send_mcp_tools_list_request(conn: "ConnectionHandler"):
async def send_mcp_tools_list_request(
conn: "ConnectionHandler", transport=None
):
"""发送MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -278,10 +219,12 @@ async def send_mcp_tools_list_request(conn: "ConnectionHandler"):
"method": "tools/list",
}
logger.bind(tag=TAG).debug("发送MCP工具列表请求")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def send_mcp_tools_list_continue_request(conn: "ConnectionHandler", cursor: str):
async def send_mcp_tools_list_continue_request(
conn: "ConnectionHandler", cursor: str, transport=None
):
"""发送带有cursor的MCP工具列表请求"""
payload = {
"jsonrpc": "2.0",
@@ -290,7 +233,7 @@ async def send_mcp_tools_list_continue_request(conn: "ConnectionHandler", cursor
"params": {"cursor": cursor},
}
logger.bind(tag=TAG).info(f"发送带cursor的MCP工具列表请求: {cursor}")
await send_mcp_message(conn, payload)
await send_mcp_message(conn, payload, transport)
async def call_mcp_tool(
@@ -299,6 +242,7 @@ async def call_mcp_tool(
tool_name: str,
args: str = "{}",
timeout: int = 30,
return_raw: bool = False,
):
"""
调用指定的工具,并等待响应
@@ -359,6 +303,7 @@ async def call_mcp_tool(
raise ValueError(f"参数必须是字典类型,实际类型: {type(arguments)}")
except Exception as e:
await mcp_client.cleanup_call_result(tool_call_id)
if not isinstance(e, ValueError):
raise ValueError(f"参数处理失败: {str(e)}")
raise e
@@ -372,7 +317,11 @@ async def call_mcp_tool(
}
logger.bind(tag=TAG).info(f"发送客户端mcp工具调用请求: {actual_name},参数: {args}")
await send_mcp_message(conn, payload)
try:
await send_mcp_message(conn, payload, raise_on_error=True)
except Exception:
await mcp_client.cleanup_call_result(tool_call_id)
raise
try:
# Wait for response or timeout
@@ -388,13 +337,21 @@ async def call_mcp_tool(
)
raise RuntimeError(f"工具调用错误: {error_msg}")
if return_raw:
return raw_result
content = raw_result.get("content")
if isinstance(content, list) and len(content) > 0:
if isinstance(content[0], dict) and "text" in content[0]:
# 直接返回文本内容,不进行JSON解析
return content[0]["text"]
# 如果结果不是预期的格式,将其转换为字符串
if return_raw:
return raw_result
return str(raw_result)
except asyncio.CancelledError:
await mcp_client.cleanup_call_result(tool_call_id)
raise
except asyncio.TimeoutError:
await mcp_client.cleanup_call_result(tool_call_id)
raise TimeoutError("工具调用请求超时")
@@ -22,6 +22,7 @@ class MCPEndpointClient:
self.lock = asyncio.Lock()
self._cached_available_tools = None # Cache for get_available_tools
self.websocket = None # WebSocket连接
self.listener_task = None
def has_tool(self, name: str) -> bool:
return name in self.tools
@@ -107,6 +108,18 @@ class MCPEndpointClient:
async def close(self):
"""关闭WebSocket连接"""
current_task = asyncio.current_task()
if (
self.listener_task
and self.listener_task is not current_task
and not self.listener_task.done()
):
self.listener_task.cancel()
try:
await self.listener_task
except asyncio.CancelledError:
pass
self.listener_task = None
if self.websocket:
await self.websocket.close()
self.websocket = None
@@ -16,6 +16,8 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
if not mcp_endpoint_url or "你的" in mcp_endpoint_url or mcp_endpoint_url == "null":
return None
websocket = None
mcp_client = None
try:
websocket = await websockets.connect(mcp_endpoint_url)
@@ -23,7 +25,9 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
mcp_client.set_websocket(websocket)
# 启动消息监听器
asyncio.create_task(_message_listener(mcp_client))
mcp_client.listener_task = asyncio.create_task(
_message_listener(mcp_client), name="xiaozhi-mcp-endpoint-listener"
)
# 发送初始化消息
await send_mcp_endpoint_initialize(mcp_client)
@@ -39,6 +43,10 @@ async def connect_mcp_endpoint(mcp_endpoint_url: str, conn=None) -> MCPEndpointC
except Exception as e:
logger.bind(tag=TAG).error(f"连接MCP接入点失败: {e}")
if mcp_client is not None:
await mcp_client.close()
elif websocket is not None:
await websocket.close()
return None
@@ -72,16 +72,25 @@ class ServerMCPClient:
async def cleanup(self):
"""清理MCP客户端资源"""
if not self._worker_task:
task = self._worker_task
if not task:
return
self._shutdown_evt.set()
try:
await asyncio.wait_for(self._worker_task, timeout=20)
except (asyncio.TimeoutError, Exception) as e:
await asyncio.wait_for(asyncio.shield(task), timeout=20)
except asyncio.TimeoutError:
self.logger.bind(tag=TAG).warning("服务端MCP关闭超时,取消工作任务")
task.cancel()
done, _ = await asyncio.wait({task}, timeout=5)
if task not in done:
self.logger.bind(tag=TAG).error("服务端MCP工作任务取消超时")
return
except Exception as e:
self.logger.bind(tag=TAG).error(f"服务端MCP客户端关闭错误: {e}")
finally:
self._worker_task = None
if task.done():
self._worker_task = None
def has_tool(self, name: str) -> bool:
"""检查是否包含指定工具
@@ -197,7 +206,7 @@ class ServerMCPClient:
if "API_ACCESS_TOKEN" in self.config:
headers["Authorization"] = f"Bearer {self.config['API_ACCESS_TOKEN']}"
self.logger.bind(tag=TAG).warning(f"你正在使用旧过时的配置 API_ACCESS_TOKEN ,请在.mcp_server_settings.json中将API_ACCESS_TOKEN直接设置在headers中,例如 'Authorization': 'Bearer API_ACCESS_TOKEN'")
# 根据transport类型选择不同的客户端,默认为SSE
transport_type = self.config.get("transport", "sse")