调通mcp tool

This commit is contained in:
玄凤科技
2025-03-20 11:52:37 +08:00
parent 8504c181c0
commit 2b81ebca8e
3 changed files with 53 additions and 59 deletions
+36 -16
View File
@@ -18,7 +18,7 @@ from concurrent.futures import ThreadPoolExecutor, TimeoutError
from core.handle.sendAudioHandle import sendAudioMessage from core.handle.sendAudioHandle import sendAudioMessage
from core.handle.receiveAudioHandle import handleAudioMessage from core.handle.receiveAudioHandle import handleAudioMessage
from core.handle.functionHandler import FunctionHandler from core.handle.functionHandler import FunctionHandler
from plugins_func.register import Action from plugins_func.register import Action, ActionResponse
from config.private_config import PrivateConfig from config.private_config import PrivateConfig
from core.auth import AuthMiddleware, AuthenticationError from core.auth import AuthMiddleware, AuthenticationError
from core.utils.auth_code_gen import AuthCodeGenerator from core.utils.auth_code_gen import AuthCodeGenerator
@@ -196,9 +196,9 @@ class ConnectionHandler:
if self.private_config: if self.private_config:
self.prompt = self.private_config.private_config.get("prompt", self.prompt) self.prompt = self.private_config.private_config.get("prompt", self.prompt)
self.client_ip_info = get_ip_info(self.client_ip) #self.client_ip_info = get_ip_info(self.client_ip)
self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}") #self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}")
self.prompt = self.prompt + f"\n我在:{self.client_ip_info}" #self.prompt = self.prompt + f"\n我在:{self.client_ip_info}"
self.dialogue.put(Message(role="system", content=self.prompt)) self.dialogue.put(Message(role="system", content=self.prompt))
self.func_handler = FunctionHandler(self.config) self.func_handler = FunctionHandler(self.config)
@@ -436,24 +436,14 @@ class ConnectionHandler:
"id": function_id, "id": function_id,
"arguments": function_arguments "arguments": function_arguments
} }
#result = self.func_handler.handle_llm_function_call(self, function_call_data)
#self._handle_function_result(result, function_call_data, text_index+1)
# 处理MCP工具调用 # 处理MCP工具调用
if self.mcp_manager.is_mcp_tool(function_name): if self.mcp_manager.is_mcp_tool(function_name):
try: result = self._handle_mcp_tool_call(function_call_data)
tool_result = asyncio.create_task(self.mcp_manager.execute_tool(
function_name,
function_arguments
))
self._handle_mcp_tool_result(tool_result, function_call_data, text_index+1)
except Exception as e:
self.logger.bind(tag=TAG).error(f"MCP工具调用错误: {e}")
response_message.append(f"MCP工具调用失败: {str(e)}")
else: else:
# 处理系统函数 # 处理系统函数
result = self.func_handler.handle_llm_function_call(self, function_call_data) result = self.func_handler.handle_llm_function_call(self, function_call_data)
self._handle_function_result(result, function_call_data, text_index+1) self._handle_function_result(result, function_call_data, text_index+1)
# 处理最后剩余的文本 # 处理最后剩余的文本
full_text = "".join(response_message) full_text = "".join(response_message)
@@ -475,6 +465,36 @@ class ConnectionHandler:
return True return True
def _handle_mcp_tool_call(self, function_call_data):
function_arguments = function_call_data["arguments"]
function_name = function_call_data["name"]
try:
args_dict = function_arguments
if isinstance(function_arguments, str):
try:
args_dict = json.loads(function_arguments)
except json.JSONDecodeError:
self.logger.bind(tag=TAG).error(f"无法解析 function_arguments: {function_arguments}")
return ActionResponse(action=Action.REQLLM, result="参数解析失败", response="")
tool_result = asyncio.run_coroutine_threadsafe(self.mcp_manager.execute_tool(
function_name,
args_dict
), self.loop).result()
# meta=None content=[TextContent(type='text', text='北京当前天气:\n温度: 21°C\n天气: 晴\n湿度: 6%\n风向: 西北 风\n风力等级: 5级', annotations=None)] isError=False
self.logger.bind(tag=TAG).info(f"tool_result:{tool_result}")
if tool_result is not None:
if tool_result.content is not None:
print(tool_result.content)
return ActionResponse(action=Action.REQLLM, result=tool_result.content[0].text, response="")
except Exception as e:
self.logger.bind(tag=TAG).error(f"MCP工具调用错误: {e}")
return ActionResponse(action=Action.REQLLM, result="工具调用出错", response="")
return ActionResponse(action=Action.REQLLM, result="工具调用出错", response="")
def _handle_function_result(self, result, function_call_data, text_index): def _handle_function_result(self, result, function_call_data, text_index):
if result.action == Action.RESPONSE: # 直接回复前端 if result.action == Action.RESPONSE: # 直接回复前端
text = result.response text = result.response
+12 -39
View File
@@ -1,7 +1,7 @@
import asyncio import asyncio
from typing import Optional from typing import Optional
from contextlib import AsyncExitStack from contextlib import AsyncExitStack
import os import os, shutil
from mcp import ClientSession, StdioServerParameters from mcp import ClientSession, StdioServerParameters
from mcp.client.stdio import stdio_client from mcp.client.stdio import stdio_client
@@ -19,34 +19,19 @@ class MCPClient:
self.tolls = [] self.tolls = []
async def initialize(self): async def initialize(self):
main_command = self.config["command"]
args = self.config.get("args", []) args = self.config.get("args", [])
server_script_path = main_command command = (
if args: shutil.which("npx")
# 如果第一个参数是路径,使用其目录作为工作目录 if self.config["command"] == "npx"
possible_path = args[0] else self.config["command"]
if os.path.exists(possible_path): )
server_script_path = possible_path
await self.connect_to_server(server_script_path)
env={**os.environ, **self.config["env"]}
async def connect_to_server(self, server_script_path: str):
"""Connect to an MCP server
Args:
server_script_path: Path to the server script (.py or .js)
"""
is_python = server_script_path.endswith('.py')
is_js = server_script_path.endswith('.js')
if not (is_python or is_js):
raise ValueError("Server script must be a .py or .js file")
env = self.config.get("env", {})
command = "python" if is_python else "node"
server_params = StdioServerParameters( server_params = StdioServerParameters(
command=command, command=command,
args=[server_script_path], args=args,
env=env env=env
) )
@@ -62,6 +47,9 @@ class MCPClient:
self.tools = tools self.tools = tools
self.logger.bind(tag=TAG).info(f"Connected to server with tools:{[tool.name for tool in tools]}") self.logger.bind(tag=TAG).info(f"Connected to server with tools:{[tool.name for tool in tools]}")
def has_tool(self, tool_name):
return any(tool.name == tool_name for tool in self.tools)
def get_available_tools(self): def get_available_tools(self):
available_tools = [{"type": "function", "function":{ available_tools = [{"type": "function", "function":{
"name": tool.name, "name": tool.name,
@@ -72,25 +60,10 @@ class MCPClient:
return available_tools return available_tools
async def call_tool(self, tool_name: str, tool_args: dict): async def call_tool(self, tool_name: str, tool_args: dict):
self.logger.bind(tag=TAG).info(f"MCPClient Calling tool {tool_name} with args: {tool_args}")
response = await self.session.call_tool(tool_name, tool_args) response = await self.session.call_tool(tool_name, tool_args)
return response return response
async def cleanup(self): async def cleanup(self):
"""Clean up resources""" """Clean up resources"""
await self.exit_stack.aclose() await self.exit_stack.aclose()
async def main():
if len(sys.argv) < 2:
print("Usage: python client.py <path_to_server_script>")
sys.exit(1)
client = MCPClient()
try:
await client.connect_to_server(sys.argv[1])
await client.chat_loop()
finally:
await client.cleanup()
if __name__ == "__main__":
import sys
asyncio.run(main())
+2 -1
View File
@@ -86,8 +86,9 @@ class MCPManager:
Raises: Raises:
ValueError: 工具未找到时抛出 ValueError: 工具未找到时抛出
""" """
self.logger.bind(tag=TAG).info(f"Executing tool {tool_name} with arguments: {arguments}")
for client in self.client.values(): for client in self.client.values():
if any(tool_name == tool["name"] for tool in client.tools): if client.has_tool(tool_name):
return await client.call_tool(tool_name, arguments) return await client.call_tool(tool_name, arguments)
raise ValueError(f"Tool {tool_name} not found in any MCP server") raise ValueError(f"Tool {tool_name} not found in any MCP server")