mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 04:43:56 +08:00
Merge branch 'main' into MVP
This commit is contained in:
+3
-3
@@ -5,15 +5,15 @@ WORKDIR /app
|
|||||||
|
|
||||||
COPY main/xiaozhi-server/requirements.txt .
|
COPY main/xiaozhi-server/requirements.txt .
|
||||||
|
|
||||||
# 优化apt安装
|
# 安装Python依赖
|
||||||
RUN pip install --no-cache-dir -r requirements.txt
|
RUN pip install --no-cache-dir -r requirements.txt
|
||||||
|
|
||||||
# 第三阶段:生产镜像
|
# 第二阶段:生产镜像
|
||||||
FROM python:3.10-slim
|
FROM python:3.10-slim
|
||||||
|
|
||||||
WORKDIR /opt/xiaozhi-esp32-server
|
WORKDIR /opt/xiaozhi-esp32-server
|
||||||
|
|
||||||
# 优化apt安装
|
# 安装系统依赖
|
||||||
RUN apt-get update && \
|
RUN apt-get update && \
|
||||||
apt-get install -y --no-install-recommends libopus0 ffmpeg && \
|
apt-get install -y --no-install-recommends libopus0 ffmpeg && \
|
||||||
apt-get clean && \
|
apt-get clean && \
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ manager-api 该项目基于SpringBoot框架开发。
|
|||||||
JDK 21
|
JDK 21
|
||||||
Maven 3.8+
|
Maven 3.8+
|
||||||
MySQL 8.0+
|
MySQL 8.0+
|
||||||
|
Redis 5.0+
|
||||||
Vue 3.x
|
Vue 3.x
|
||||||
|
|
||||||
# 创建数据库
|
# 创建数据库
|
||||||
@@ -43,6 +44,37 @@ spring:
|
|||||||
password: 123456
|
password: 123456
|
||||||
```
|
```
|
||||||
|
|
||||||
|
|
||||||
|
# 连接Redis
|
||||||
|
|
||||||
|
如果还没有Redis,你可以通过docker安装redis
|
||||||
|
|
||||||
|
```
|
||||||
|
docker run --name xiaozhi-esp32-server-redis -d -p 6379:6379 redis
|
||||||
|
```
|
||||||
|
|
||||||
|
# 确认项目Redis连接信息
|
||||||
|
|
||||||
|
在`src/main/resources/application-dev.yml`中配置Redis连接信息
|
||||||
|
|
||||||
|
```
|
||||||
|
spring:
|
||||||
|
data:
|
||||||
|
redis:
|
||||||
|
host: localhost
|
||||||
|
port: 6379
|
||||||
|
password:
|
||||||
|
database: 0
|
||||||
|
```
|
||||||
|
|
||||||
|
在`src/main/resources/application.yml`中配置Redis启用状态
|
||||||
|
|
||||||
|
```
|
||||||
|
renren:
|
||||||
|
redis:
|
||||||
|
open: true
|
||||||
|
```
|
||||||
|
|
||||||
# 测试启动
|
# 测试启动
|
||||||
|
|
||||||
本项目为SpringBoot项目,启动方式为:
|
本项目为SpringBoot项目,启动方式为:
|
||||||
|
|||||||
Binary file not shown.
|
After Width: | Height: | Size: 54 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 13 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 32 KiB |
@@ -39,8 +39,8 @@ export default {
|
|||||||
},
|
},
|
||||||
methods: {
|
methods: {
|
||||||
confirm() {
|
confirm() {
|
||||||
if (!this.deviceCode.trim()) {
|
if (!/^\d{6}$/.test(this.deviceCode)) {
|
||||||
this.$message.error('请输入设备验证码');
|
this.$message.error('请输入6位数字设备验证码');
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
this.$emit('confirm',this.deviceCode)
|
this.$emit('confirm',this.deviceCode)
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ import sys
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
from config.settings import load_config
|
from config.settings import load_config
|
||||||
|
|
||||||
SERVER_VERSION = "0.1.15"
|
SERVER_VERSION = "0.1.16"
|
||||||
|
|
||||||
def setup_logging():
|
def setup_logging():
|
||||||
"""从配置文件中读取日志配置,并设置日志输出格式和级别"""
|
"""从配置文件中读取日志配置,并设置日志输出格式和级别"""
|
||||||
@@ -13,7 +13,7 @@ def setup_logging():
|
|||||||
log_format_file = log_config.get("log_format_file", "{time:YYYY-MM-DD HH:mm:ss} - {version_{selected_module}} - {name} - {level} - {extra[tag]} - {message}")
|
log_format_file = log_config.get("log_format_file", "{time:YYYY-MM-DD HH:mm:ss} - {version_{selected_module}} - {name} - {level} - {extra[tag]} - {message}")
|
||||||
|
|
||||||
selected_module = config.get("selected_module")
|
selected_module = config.get("selected_module")
|
||||||
selected_module_str = ''.join([key[0] + value[0] for key, value in selected_module.items()])
|
selected_module_str = ''.join([value[0] + value[1] for key, value in selected_module.items()])
|
||||||
|
|
||||||
log_format = log_format.replace("{version}", SERVER_VERSION)
|
log_format = log_format.replace("{version}", SERVER_VERSION)
|
||||||
log_format = log_format.replace("{selected_module}", selected_module_str)
|
log_format = log_format.replace("{selected_module}", selected_module_str)
|
||||||
|
|||||||
@@ -18,10 +18,11 @@ 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
|
||||||
|
from core.mcp.manager import MCPManager
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -101,6 +102,8 @@ class ConnectionHandler:
|
|||||||
if self.config["selected_module"]["Intent"] == 'function_call':
|
if self.config["selected_module"]["Intent"] == 'function_call':
|
||||||
self.use_function_call_mode = True
|
self.use_function_call_mode = True
|
||||||
|
|
||||||
|
self.mcp_manager = MCPManager(self)
|
||||||
|
|
||||||
async def handle_connection(self, ws):
|
async def handle_connection(self, ws):
|
||||||
try:
|
try:
|
||||||
# 获取并验证headers
|
# 获取并验证headers
|
||||||
@@ -194,7 +197,31 @@ class ConnectionHandler:
|
|||||||
"""加载记忆"""
|
"""加载记忆"""
|
||||||
device_id = self.headers.get("device-id", None)
|
device_id = self.headers.get("device-id", None)
|
||||||
self.memory.init_memory(device_id, self.llm)
|
self.memory.init_memory(device_id, self.llm)
|
||||||
self.intent.set_llm(self.llm)
|
|
||||||
|
"""为意图识别设置LLM,优先使用专用LLM"""
|
||||||
|
# 检查是否配置了专用的意图识别LLM
|
||||||
|
intent_llm_name = self.config["Intent"]["intent_llm"]["llm"]
|
||||||
|
|
||||||
|
# 记录开始初始化意图识别LLM的时间
|
||||||
|
intent_llm_init_start = time.time()
|
||||||
|
|
||||||
|
if not self.use_function_call_mode and intent_llm_name and intent_llm_name in self.config["LLM"]:
|
||||||
|
# 如果配置了专用LLM,则创建独立的LLM实例
|
||||||
|
from core.utils import llm as llm_utils
|
||||||
|
intent_llm_config = self.config["LLM"][intent_llm_name]
|
||||||
|
intent_llm_type = intent_llm_config.get("type", intent_llm_name)
|
||||||
|
intent_llm = llm_utils.create_instance(intent_llm_type, intent_llm_config)
|
||||||
|
self.logger.bind(tag=TAG).info(f"为意图识别创建了专用LLM: {intent_llm_name}, 类型: {intent_llm_type}")
|
||||||
|
|
||||||
|
self.intent.set_llm(intent_llm)
|
||||||
|
else:
|
||||||
|
# 否则使用主LLM
|
||||||
|
self.intent.set_llm(self.llm)
|
||||||
|
self.logger.bind(tag=TAG).info("意图识别使用主LLM")
|
||||||
|
|
||||||
|
# 记录意图识别LLM初始化耗时
|
||||||
|
intent_llm_init_time = time.time() - intent_llm_init_start
|
||||||
|
self.logger.bind(tag=TAG).info(f"意图识别LLM初始化完成,耗时: {intent_llm_init_time:.4f}秒")
|
||||||
|
|
||||||
"""加载位置信息"""
|
"""加载位置信息"""
|
||||||
self.client_ip_info = get_ip_info(self.client_ip)
|
self.client_ip_info = get_ip_info(self.client_ip)
|
||||||
@@ -203,6 +230,9 @@ class ConnectionHandler:
|
|||||||
self.prompt = self.prompt + f"\nuser location:{self.client_ip_info}"
|
self.prompt = self.prompt + f"\nuser location:{self.client_ip_info}"
|
||||||
self.dialogue.update_system_message(self.prompt)
|
self.dialogue.update_system_message(self.prompt)
|
||||||
|
|
||||||
|
"""加载MCP工具"""
|
||||||
|
asyncio.run_coroutine_threadsafe(self.mcp_manager.initialize_servers(), self.loop)
|
||||||
|
|
||||||
def change_system_prompt(self, prompt):
|
def change_system_prompt(self, prompt):
|
||||||
self.prompt = prompt
|
self.prompt = prompt
|
||||||
# 找到原来的role==system,替换原来的系统提示
|
# 找到原来的role==system,替换原来的系统提示
|
||||||
@@ -324,7 +354,6 @@ class ConnectionHandler:
|
|||||||
functions = None
|
functions = None
|
||||||
if hasattr(self, 'func_handler'):
|
if hasattr(self, 'func_handler'):
|
||||||
functions = self.func_handler.get_functions()
|
functions = self.func_handler.get_functions()
|
||||||
|
|
||||||
response_message = []
|
response_message = []
|
||||||
processed_chars = 0 # 跟踪已处理的字符位置
|
processed_chars = 0 # 跟踪已处理的字符位置
|
||||||
|
|
||||||
@@ -358,6 +387,9 @@ class ConnectionHandler:
|
|||||||
content_arguments = ""
|
content_arguments = ""
|
||||||
for response in llm_responses:
|
for response in llm_responses:
|
||||||
content, tools_call = response
|
content, tools_call = response
|
||||||
|
if "content" in response:
|
||||||
|
content = response["content"]
|
||||||
|
tools_call = None
|
||||||
if content is not None and len(content) > 0:
|
if content is not None and len(content) > 0:
|
||||||
if len(response_message) <= 0 and (content == "```" or "<tool_call>" in content):
|
if len(response_message) <= 0 and (content == "```" or "<tool_call>" in content):
|
||||||
tool_call_flag = True
|
tool_call_flag = True
|
||||||
@@ -436,8 +468,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工具调用
|
||||||
|
if self.mcp_manager.is_mcp_tool(function_name):
|
||||||
|
result = self._handle_mcp_tool_call(function_call_data)
|
||||||
|
else:
|
||||||
|
# 处理系统函数
|
||||||
|
result = self.func_handler.handle_llm_function_call(self, function_call_data)
|
||||||
|
self._handle_function_result(result, function_call_data, text_index+1)
|
||||||
|
|
||||||
# 处理最后剩余的文本
|
# 处理最后剩余的文本
|
||||||
full_text = "".join(response_message)
|
full_text = "".join(response_message)
|
||||||
@@ -459,6 +497,42 @@ 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
|
||||||
|
content_text = ""
|
||||||
|
if tool_result is not None and tool_result.content is not None:
|
||||||
|
for content in tool_result.content:
|
||||||
|
content_type = content.type
|
||||||
|
if content_type == "text":
|
||||||
|
content_text = content.text
|
||||||
|
elif content_type == "image":
|
||||||
|
pass
|
||||||
|
|
||||||
|
if len(content_text) > 0:
|
||||||
|
return ActionResponse(action=Action.REQLLM, result=content_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
|
||||||
@@ -582,6 +656,8 @@ class ConnectionHandler:
|
|||||||
|
|
||||||
async def close(self, ws=None):
|
async def close(self, ws=None):
|
||||||
"""资源清理方法"""
|
"""资源清理方法"""
|
||||||
|
# 清理MCP资源
|
||||||
|
await self.mcp_manager.cleanup_all()
|
||||||
|
|
||||||
# 触发停止事件并清理资源
|
# 触发停止事件并清理资源
|
||||||
if self.stop_event:
|
if self.stop_event:
|
||||||
|
|||||||
@@ -4,6 +4,8 @@ import uuid
|
|||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
from core.handle.helloHandle import checkWakeupWords
|
from core.handle.helloHandle import checkWakeupWords
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
|
from core.utils.dialogue import Message
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -21,11 +23,11 @@ async def handle_user_intent(conn, text):
|
|||||||
# 使用支持function calling的聊天方法,不再进行意图分析
|
# 使用支持function calling的聊天方法,不再进行意图分析
|
||||||
return False
|
return False
|
||||||
# 使用LLM进行意图分析
|
# 使用LLM进行意图分析
|
||||||
intent = await analyze_intent_with_llm(conn, text)
|
intent_result = await analyze_intent_with_llm(conn, text)
|
||||||
if not intent:
|
if not intent_result:
|
||||||
return False
|
return False
|
||||||
# 处理各种意图
|
# 处理各种意图
|
||||||
return await process_intent_result(conn, intent, text)
|
return await process_intent_result(conn, intent_result, text)
|
||||||
|
|
||||||
|
|
||||||
async def check_direct_exit(conn, text):
|
async def check_direct_exit(conn, text):
|
||||||
@@ -40,7 +42,6 @@ async def check_direct_exit(conn, text):
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
async def analyze_intent_with_llm(conn, text):
|
async def analyze_intent_with_llm(conn, text):
|
||||||
"""使用LLM分析用户意图"""
|
"""使用LLM分析用户意图"""
|
||||||
if not hasattr(conn, 'intent') or not conn.intent:
|
if not hasattr(conn, 'intent') or not conn.intent:
|
||||||
@@ -51,52 +52,65 @@ async def analyze_intent_with_llm(conn, text):
|
|||||||
dialogue = conn.dialogue
|
dialogue = conn.dialogue
|
||||||
try:
|
try:
|
||||||
intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text)
|
intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text)
|
||||||
# 尝试解析JSON结果
|
return intent_result
|
||||||
try:
|
|
||||||
intent_data = json.loads(intent_result)
|
|
||||||
if "intent" in intent_data:
|
|
||||||
return intent_data["intent"]
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
# 如果不是JSON格式,尝试直接获取意图文本
|
|
||||||
return intent_result.strip()
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}")
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
async def process_intent_result(conn, intent, original_text):
|
async def process_intent_result(conn, intent_result, original_text):
|
||||||
"""处理意图识别结果"""
|
"""处理意图识别结果"""
|
||||||
# 处理退出意图
|
try:
|
||||||
if "结束聊天" in intent:
|
# 尝试将结果解析为JSON
|
||||||
logger.bind(tag=TAG).info(f"识别到退出意图: {intent}")
|
intent_data = json.loads(intent_result)
|
||||||
# 如果是明确的离别意图,发送告别语并关闭连接
|
|
||||||
await send_stt_message(conn, original_text)
|
|
||||||
conn.executor.submit(conn.chat_and_close, original_text)
|
|
||||||
return True
|
|
||||||
|
|
||||||
# 处理播放音乐意图
|
# 检查是否有function_call
|
||||||
if "播放音乐" in intent:
|
if "function_call" in intent_data:
|
||||||
logger.bind(tag=TAG).info(f"识别到音乐播放意图: {intent}")
|
# 直接从意图识别获取了function_call
|
||||||
# 调用play_music函数来播放音乐
|
logger.bind(tag=TAG).info(f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}")
|
||||||
song_name = extract_text_in_brackets(intent)
|
function_name = intent_data["function_call"]["name"]
|
||||||
function_id = str(uuid.uuid4().hex)
|
if function_name == "continue_chat":
|
||||||
function_name = "play_music"
|
return False
|
||||||
function_arguments = '{ "song_name": "' + song_name + '" }'
|
function_args = None
|
||||||
|
if "arguments" in intent_data["function_call"]:
|
||||||
|
function_args = intent_data["function_call"]["arguments"]
|
||||||
|
# 确保参数是字符串格式的JSON
|
||||||
|
if isinstance(function_args, dict):
|
||||||
|
function_args = json.dumps(function_args)
|
||||||
|
|
||||||
function_call_data = {
|
function_call_data = {
|
||||||
"name": function_name,
|
"name": function_name,
|
||||||
"id": function_id,
|
"id": str(uuid.uuid4().hex),
|
||||||
"arguments": function_arguments
|
"arguments": function_args
|
||||||
}
|
}
|
||||||
conn.func_handler.handle_llm_function_call(conn, function_call_data)
|
|
||||||
return True
|
|
||||||
|
|
||||||
# 其他意图处理可以在这里扩展
|
await send_stt_message(conn, original_text)
|
||||||
|
|
||||||
# 默认返回False,表示继续常规聊天流程
|
# 使用executor执行函数调用和结果处理
|
||||||
return False
|
def process_function_call():
|
||||||
|
conn.dialogue.put(Message(role="user", content=original_text))
|
||||||
|
result = conn.func_handler.handle_llm_function_call(conn, function_call_data)
|
||||||
|
if result and function_name != 'play_music':
|
||||||
|
# 获取当前最新的文本索引
|
||||||
|
text = result.response
|
||||||
|
if text is None:
|
||||||
|
text = result.result
|
||||||
|
if text is not None:
|
||||||
|
text_index = conn.tts_last_text_index + 1 if hasattr(conn, 'tts_last_text_index') else 0
|
||||||
|
conn.recode_first_last_text(text, text_index)
|
||||||
|
future = conn.executor.submit(conn.speak_and_play, text, text_index)
|
||||||
|
conn.llm_finish_task = True
|
||||||
|
conn.tts_queue.put(future)
|
||||||
|
conn.dialogue.put(Message(role="assistant", content=text))
|
||||||
|
|
||||||
|
# 将函数执行放在线程池中
|
||||||
|
conn.executor.submit(process_function_call)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
except json.JSONDecodeError as e:
|
||||||
|
logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}")
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def extract_text_in_brackets(s):
|
def extract_text_in_brackets(s):
|
||||||
|
|||||||
@@ -0,0 +1,86 @@
|
|||||||
|
from datetime import timedelta
|
||||||
|
from typing import Optional
|
||||||
|
from contextlib import AsyncExitStack
|
||||||
|
import os, shutil
|
||||||
|
from mcp import ClientSession, StdioServerParameters
|
||||||
|
from mcp.client.stdio import stdio_client
|
||||||
|
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
|
||||||
|
class MCPClient:
|
||||||
|
def __init__(self, config):
|
||||||
|
# Initialize session and client objects
|
||||||
|
self.session: Optional[ClientSession] = None
|
||||||
|
self.exit_stack = AsyncExitStack()
|
||||||
|
self.logger = setup_logging()
|
||||||
|
self.config = config
|
||||||
|
self.tolls = []
|
||||||
|
|
||||||
|
async def initialize(self):
|
||||||
|
args = self.config.get("args", [])
|
||||||
|
|
||||||
|
command = (
|
||||||
|
shutil.which("npx")
|
||||||
|
if self.config["command"] == "npx"
|
||||||
|
else self.config["command"]
|
||||||
|
)
|
||||||
|
|
||||||
|
env={**os.environ}
|
||||||
|
if self.config.get("env"):
|
||||||
|
env.update(self.config["env"])
|
||||||
|
|
||||||
|
server_params = StdioServerParameters(
|
||||||
|
command=command,
|
||||||
|
args=args,
|
||||||
|
env=env
|
||||||
|
)
|
||||||
|
|
||||||
|
stdio_transport = await self.exit_stack.enter_async_context(stdio_client(server_params))
|
||||||
|
self.stdio, self.write = stdio_transport
|
||||||
|
time_out_delta = timedelta(seconds=15)
|
||||||
|
self.session = await self.exit_stack.enter_async_context(ClientSession(read_stream=self.stdio, write_stream=self.write, read_timeout_seconds=time_out_delta))
|
||||||
|
|
||||||
|
await self.session.initialize()
|
||||||
|
|
||||||
|
# List available tools
|
||||||
|
response = await self.session.list_tools()
|
||||||
|
tools = response.tools
|
||||||
|
self.tools = 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):
|
||||||
|
available_tools = [{"type": "function", "function":{
|
||||||
|
"name": tool.name,
|
||||||
|
"description": tool.description,
|
||||||
|
"parameters": tool.inputSchema
|
||||||
|
} } for tool in self.tools]
|
||||||
|
|
||||||
|
return available_tools
|
||||||
|
|
||||||
|
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}")
|
||||||
|
try:
|
||||||
|
response = await self.session.call_tool(tool_name, tool_args)
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"Error calling tool {tool_name}: {e}")
|
||||||
|
from types import SimpleNamespace
|
||||||
|
error_content = SimpleNamespace(
|
||||||
|
type='text',
|
||||||
|
text=f"Error calling tool {tool_name}: {e}"
|
||||||
|
)
|
||||||
|
error_response = SimpleNamespace(
|
||||||
|
content=[error_content],
|
||||||
|
isError=True
|
||||||
|
)
|
||||||
|
return error_response
|
||||||
|
self.logger.bind(tag=TAG).info(f"MCPClient Response from tool {tool_name}: {response}")
|
||||||
|
return response
|
||||||
|
|
||||||
|
async def cleanup(self):
|
||||||
|
"""Clean up resources"""
|
||||||
|
await self.exit_stack.aclose()
|
||||||
@@ -0,0 +1,110 @@
|
|||||||
|
"""MCP服务管理器"""
|
||||||
|
import os, json
|
||||||
|
from typing import Dict, Any, List
|
||||||
|
from .MCPClient import MCPClient
|
||||||
|
from config.logger import setup_logging
|
||||||
|
from core.utils.util import get_project_dir
|
||||||
|
from plugins_func.register import register_function, ActionResponse, Action, ToolType
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
|
||||||
|
class MCPManager:
|
||||||
|
"""管理多个MCP服务的集中管理器"""
|
||||||
|
|
||||||
|
def __init__(self,conn) -> None:
|
||||||
|
"""
|
||||||
|
初始化MCP管理器
|
||||||
|
"""
|
||||||
|
self.conn = conn
|
||||||
|
self.logger = setup_logging()
|
||||||
|
self.config_path = get_project_dir() + 'data/.mcp_server_settings.json'
|
||||||
|
if os.path.exists(self.config_path) == False:
|
||||||
|
self.config_path = ""
|
||||||
|
self.logger.bind(tag=TAG).warning(f"请检查mcp服务配置文件:data/.mcp_server_settings.json")
|
||||||
|
self.client: Dict[str, MCPClient] = {}
|
||||||
|
self.tools = []
|
||||||
|
|
||||||
|
def load_config(self) -> Dict[str, Any]:
|
||||||
|
"""加载MCP服务配置
|
||||||
|
Returns:
|
||||||
|
Dict[str, Any]: 服务配置字典
|
||||||
|
"""
|
||||||
|
if len(self.config_path) == 0:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
try:
|
||||||
|
with open(self.config_path, 'r', encoding='utf-8') as f:
|
||||||
|
config = json.load(f)
|
||||||
|
return config.get('mcpServers', {})
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"Error loading MCP config from {self.config_path}: {e}")
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def initialize_servers(self) -> None:
|
||||||
|
"""初始化所有MCP服务"""
|
||||||
|
config = self.load_config()
|
||||||
|
for name, srv_config in config.items():
|
||||||
|
if not srv_config.get("command"):
|
||||||
|
self.logger.bind(tag=TAG).warning(f"Skipping server {name}: command not specified")
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
client = MCPClient(srv_config)
|
||||||
|
await client.initialize()
|
||||||
|
self.client[name] = client
|
||||||
|
self.logger.bind(tag=TAG).info(f"Initialized MCP client: {name}")
|
||||||
|
client_tools = client.get_available_tools()
|
||||||
|
self.tools.extend(client_tools)
|
||||||
|
for tool in client_tools:
|
||||||
|
func_name = "mcp_"+tool["function"]["name"]
|
||||||
|
register_function(func_name, tool, ToolType.MCP_CLIENT)(self.execute_tool)
|
||||||
|
self.conn.func_handler.function_registry.register_function(func_name)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"Failed to initialize MCP server {name}: {e}")
|
||||||
|
self.conn.func_handler.upload_functions_desc()
|
||||||
|
|
||||||
|
def get_all_tools(self) -> List[Dict[str, Any]]:
|
||||||
|
"""获取所有服务的工具function定义
|
||||||
|
Returns:
|
||||||
|
List[Dict[str, Any]]: 所有工具的function定义列表
|
||||||
|
"""
|
||||||
|
return self.tools
|
||||||
|
|
||||||
|
def is_mcp_tool(self, tool_name: str) -> bool:
|
||||||
|
"""检查是否是MCP工具
|
||||||
|
Args:
|
||||||
|
tool_name: 工具名称
|
||||||
|
Returns:
|
||||||
|
bool: 是否是MCP工具
|
||||||
|
"""
|
||||||
|
for tool in self.tools:
|
||||||
|
if tool.get("function") != None and tool["function"].get("name") == tool_name:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def execute_tool(self, tool_name: str, arguments: Dict[str, Any]) -> Any:
|
||||||
|
"""执行工具调用
|
||||||
|
Args:
|
||||||
|
tool_name: 工具名称
|
||||||
|
arguments: 工具参数
|
||||||
|
Returns:
|
||||||
|
Any: 工具执行结果
|
||||||
|
Raises:
|
||||||
|
ValueError: 工具未找到时抛出
|
||||||
|
"""
|
||||||
|
self.logger.bind(tag=TAG).info(f"Executing tool {tool_name} with arguments: {arguments}")
|
||||||
|
for client in self.client.values():
|
||||||
|
if client.has_tool(tool_name):
|
||||||
|
return await client.call_tool(tool_name, arguments)
|
||||||
|
|
||||||
|
raise ValueError(f"Tool {tool_name} not found in any MCP server")
|
||||||
|
|
||||||
|
async def cleanup_all(self) -> None:
|
||||||
|
for name, client in self.client.items():
|
||||||
|
try:
|
||||||
|
await client.cleanup()
|
||||||
|
self.logger.bind(tag=TAG).info(f"Cleaned up MCP client: {name}")
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"Error cleaning up MCP client {name}: {e}")
|
||||||
|
self.client.clear()
|
||||||
@@ -10,14 +10,21 @@ class IntentProviderBase(ABC):
|
|||||||
def __init__(self, config):
|
def __init__(self, config):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.intent_options = config.get("intent_options", {
|
self.intent_options = config.get("intent_options", {
|
||||||
|
"handle_exit_intent": "结束聊天, 用户发来如再见之类的表示结束的话, 不想再进行对话的时候",
|
||||||
|
"play_music": "播放音乐, 用户希望你可以播放音乐, 只用于播放音乐的意图",
|
||||||
|
"get_weather": "查询天气, 用户希望查询某个地点的天气情况",
|
||||||
|
"get_news": "查询新闻, 用户希望查询最新新闻或特定类型的新闻",
|
||||||
|
"get_lunar": "用于获取今天的阴历/农历和黄历信息",
|
||||||
|
"get_time": "获取今天日期或者当前时间信息",
|
||||||
"continue_chat": "继续聊天",
|
"continue_chat": "继续聊天",
|
||||||
"end_chat": "结束聊天",
|
|
||||||
"play_music": "播放音乐"
|
|
||||||
})
|
})
|
||||||
|
|
||||||
def set_llm(self, llm):
|
def set_llm(self, llm):
|
||||||
self.llm = llm
|
self.llm = llm
|
||||||
logger.bind(tag=TAG).debug("Set LLM for intent provider")
|
# 获取模型名称和类型信息
|
||||||
|
model_name = getattr(llm, 'model_name', str(llm.__class__.__name__))
|
||||||
|
# 记录更详细的日志
|
||||||
|
logger.bind(tag=TAG).info(f"意图识别设置LLM: {model_name}")
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str:
|
async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str:
|
||||||
@@ -30,5 +37,6 @@ class IntentProviderBase(ABC):
|
|||||||
- "继续聊天"
|
- "继续聊天"
|
||||||
- "结束聊天"
|
- "结束聊天"
|
||||||
- "播放音乐 歌名" 或 "随机播放音乐"
|
- "播放音乐 歌名" 或 "随机播放音乐"
|
||||||
|
- "查询天气 地点名" 或 "查询天气 [当前位置]"
|
||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
from ..base import IntentProviderBase
|
||||||
|
from typing import List, Dict
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
class IntentProvider(IntentProviderBase):
|
||||||
|
async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str:
|
||||||
|
"""
|
||||||
|
默认的意图识别实现,始终返回继续聊天
|
||||||
|
Args:
|
||||||
|
dialogue_history: 对话历史记录列表
|
||||||
|
text: 本次对话记录
|
||||||
|
Returns:
|
||||||
|
固定返回"继续聊天"
|
||||||
|
"""
|
||||||
|
logger.bind(tag=TAG).debug("Using functionCallProvider, always returning continue chat")
|
||||||
|
return self.intent_options["continue_chat"]
|
||||||
@@ -3,6 +3,9 @@ from ..base import IntentProviderBase
|
|||||||
from plugins_func.functions.play_music import initialize_music_handler
|
from plugins_func.functions.play_music import initialize_music_handler
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
import re
|
import re
|
||||||
|
import json
|
||||||
|
import hashlib
|
||||||
|
import time
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -13,6 +16,10 @@ class IntentProvider(IntentProviderBase):
|
|||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.llm = None
|
self.llm = None
|
||||||
self.promot = self.get_intent_system_prompt()
|
self.promot = self.get_intent_system_prompt()
|
||||||
|
# 添加缓存管理
|
||||||
|
self.intent_cache = {} # 缓存意图识别结果
|
||||||
|
self.cache_expiry = 600 # 缓存有效期10分钟
|
||||||
|
self.cache_max_size = 100 # 最多缓存100个意图
|
||||||
|
|
||||||
def get_intent_system_prompt(self) -> str:
|
def get_intent_system_prompt(self) -> str:
|
||||||
"""
|
"""
|
||||||
@@ -22,63 +29,120 @@ class IntentProvider(IntentProviderBase):
|
|||||||
"""
|
"""
|
||||||
intent_list = []
|
intent_list = []
|
||||||
|
|
||||||
"""
|
|
||||||
"continue_chat": "1.继续聊天, 除了播放音乐和结束聊天的时候的选项, 比如日常的聊天和问候, 对话等",
|
|
||||||
"end_chat": "2.结束聊天, 用户发来如再见之类的表示结束的话, 不想再进行对话的时候",
|
|
||||||
"play_music": "3.播放音乐, 用户希望你可以播放音乐, 只用于播放音乐的意图"
|
|
||||||
"""
|
|
||||||
for key, value in self.intent_options.items():
|
|
||||||
if key == "play_music":
|
|
||||||
intent_list.append("3.播放音乐, 用户希望你可以播放音乐, 只用于播放音乐的意图")
|
|
||||||
elif key == "end_chat":
|
|
||||||
intent_list.append("2.结束聊天, 用户发来如再见之类的表示结束的话, 不想再进行对话的时候")
|
|
||||||
elif key == "continue_chat":
|
|
||||||
intent_list.append("1.继续聊天, 除了播放音乐和结束聊天的时候的选项, 比如日常的聊天和问候, 对话等")
|
|
||||||
else:
|
|
||||||
intent_list.append(value)
|
|
||||||
|
|
||||||
# "如果是唱歌、听歌、播放音乐,请指定歌名,格式为'播放音乐 [识别出的歌名]'。\n"
|
|
||||||
# "如果听不出具体歌名,可以返回'随机播放音乐'。\n"
|
|
||||||
# "只需要返回意图结果的json,不要解释。"
|
|
||||||
# "返回格式如下:\n"
|
|
||||||
prompt = (
|
prompt = (
|
||||||
"你是一个意图识别助手。你需要根据和用户的对话记录,重点分析用户的最后一句话,判断用户意图属于以下哪一类(使用<start>和<end>标志):\n"
|
"你是一个意图识别助手。请分析用户的最后一句话,判断用户意图属于以下哪一类:\n"
|
||||||
"<start>"
|
"<start>"
|
||||||
f"{', '.join(intent_list)}"
|
f"{', '.join(intent_list)}"
|
||||||
"<end>\n"
|
"<end>\n"
|
||||||
"你需要按照以下的步骤处理用户的对话"
|
"处理步骤:"
|
||||||
"1. 思考出对话的意图是哪一类的"
|
"1. 思考意图类型,生成function_call格式"
|
||||||
"2. 属于1和2的意图, 直接返回,返回格式如下:\n"
|
"\n\n"
|
||||||
"{intent: '用户意图'}\n"
|
"返回格式示例:\n"
|
||||||
"3. 属于3的意图,则继续分析用户希望播放的音乐\n"
|
"1. 播放音乐意图: {\"function_call\": {\"name\": \"play_music\", \"arguments\": {\"song_name\": \"音乐名称\"}}}\n"
|
||||||
"4. 如果无法识别出具体歌名,可以返回'随机播放音乐'\n"
|
"2. 查询天气意图: {\"function_call\": {\"name\": \"get_weather\", \"arguments\": {\"location\": \"地点名称\", \"lang\": \"zh_CN\"}}}\n"
|
||||||
"{intent: '播放音乐 [获取的音乐名字]'}\n"
|
"3. 查询新闻意图: {\"function_call\": {\"name\": \"get_news\", \"arguments\": {\"category\": \"新闻类别\", \"detail\": false, \"lang\": \"zh_CN\"}}}\n"
|
||||||
"下面是几个处理的示例(思考的内容不返回, 只返回json部分, 无额外的内容)\n"
|
"4. 结束对话意图: {\"function_call\": {\"name\": \"handle_exit_intent\", \"arguments\": {\"say_goodbye\": \"goodbye\"}}}\n"
|
||||||
"```"
|
"5. 获取当天日期时间: {\"function_call\": {\"name\": \"get_time\"}}\n"
|
||||||
|
"6. 获取当前黄历意图: {\"function_call\": {\"name\": \"get_lunar\"}}\n"
|
||||||
|
"7. 继续聊天意图: {\"function_call\": {\"name\": \"continue_chat\"}}\n"
|
||||||
|
"\n"
|
||||||
|
"注意:\n"
|
||||||
|
"- 播放音乐:无歌名时,song_name设为\"random\"\n"
|
||||||
|
"- 查询天气:无地点时,location设为null\n"
|
||||||
|
"- 查询新闻:无类别时,category设为null;查询详情时,detail设为true\n"
|
||||||
|
"- 如果没有明显的意图,应按照继续聊天意图处理\n"
|
||||||
|
"- 只返回纯JSON,不要任何其他内容\n"
|
||||||
|
"\n"
|
||||||
|
"示例分析:\n"
|
||||||
|
"```\n"
|
||||||
|
"用户: 你好小智\n"
|
||||||
|
"返回: {\"function_call\": {\"name\": \"continue_chat\"}}\n"
|
||||||
|
"```\n"
|
||||||
|
"```\n"
|
||||||
"用户: 你今天怎么样?\n"
|
"用户: 你今天怎么样?\n"
|
||||||
"思考(不返回): 用户发来的数据是一个问候语,属于继续聊天的意图, 是种类1, 种类1的需求是直接返回\n"
|
"返回: {\"function_call\": {\"name\": \"continue_chat\"}}\n"
|
||||||
"返回结果: {intent: '继续聊天'}\n"
|
"```\n"
|
||||||
"```"
|
"```\n"
|
||||||
"用户: 我今天有点累了, 我们明天再聊吧\n"
|
"用户: 现在是几号了?现在几点了?\n"
|
||||||
"思考(不返回): 用户表达了今天不想继续对话,属于结束聊天的意图, 是种类2, 种类2的需求是直接返回\n"
|
"返回: {\"function_call\": {\"name\": \"get_time\"}}\n"
|
||||||
"返回结果: {intent: '结束聊天'}\n"
|
"```\n"
|
||||||
"```"
|
"```\n"
|
||||||
"用户: 我今天有点累了, 我们明天再聊吧\n"
|
"用户: 今天农历是多少?\n"
|
||||||
"思考(不返回): 用户表达了今天不想继续对话,属于结束聊天的意图, 是种类2, 种类2的需求是直接返回\n"
|
"返回: {\"function_call\": {\"name\": \"get_lunar\"}}\n"
|
||||||
"返回结果: {intent: '结束聊天'}\n"
|
"```\n"
|
||||||
"```"
|
"```\n"
|
||||||
"用户: 你可以播放一首中秋月给我听吗\n"
|
"用户: 我们明天再聊吧\n"
|
||||||
"思考(不返回): 用户表达了想听音乐的续签,属于播放音乐的意图, 是种类3, 种类3的需求需要继续判断播放的音乐, 这里用户希望的歌曲名明确给出是中秋月\n"
|
"返回: {\"function_call\": {\"name\": \"handle_exit_intent\"}}\n"
|
||||||
"返回结果: {intent: '播放音乐 [中秋月]'}\n"
|
"```\n"
|
||||||
"```"
|
"```\n"
|
||||||
"你现在可以使用的音乐的名称如下(使用<start>和<end>标志):\n"
|
"用户: 播放中秋月\n"
|
||||||
|
"返回: {\"function_call\": {\"name\": \"play_music\", \"arguments\": {\"song_name\": \"中秋月\"}}}\n"
|
||||||
|
"```\n"
|
||||||
|
"```\n"
|
||||||
|
"用户: 北京天气怎么样\n"
|
||||||
|
"返回: {\"function_call\": {\"name\": \"get_weather\", \"arguments\": {\"location\": \"北京\", \"lang\": \"zh_CN\"}}}\n"
|
||||||
|
"```\n"
|
||||||
|
"```\n"
|
||||||
|
"用户: 今天天气怎么样\n"
|
||||||
|
"返回: {\"function_call\": {\"name\": \"get_weather\", \"arguments\": {\"location\": null, \"lang\": \"zh_CN\"}}}\n"
|
||||||
|
"```\n"
|
||||||
|
"```\n"
|
||||||
|
"用户: 播报财经新闻\n"
|
||||||
|
"返回: {\"function_call\": {\"name\": \"get_news\", \"arguments\": {\"category\": \"财经\", \"detail\": false, \"lang\": \"zh_CN\"}}}\n"
|
||||||
|
"```\n"
|
||||||
|
"```\n"
|
||||||
|
"用户: 有什么最新新闻\n"
|
||||||
|
"返回: {\"function_call\": {\"name\": \"get_news\", \"arguments\": {\"category\": null, \"detail\": false, \"lang\": \"zh_CN\"}}}\n"
|
||||||
|
"```\n"
|
||||||
|
"```\n"
|
||||||
|
"用户: 详细介绍一下这条新闻\n"
|
||||||
|
"返回: {\"function_call\": {\"name\": \"get_news\", \"arguments\": {\"detail\": true, \"lang\": \"zh_CN\"}}}\n"
|
||||||
|
"```\n"
|
||||||
|
"可用的音乐名称:\n"
|
||||||
)
|
)
|
||||||
return prompt
|
return prompt
|
||||||
|
|
||||||
|
def clean_cache(self):
|
||||||
|
"""清理过期缓存"""
|
||||||
|
now = time.time()
|
||||||
|
# 找出过期键
|
||||||
|
expired_keys = [k for k, v in self.intent_cache.items() if now - v['timestamp'] > self.cache_expiry]
|
||||||
|
for key in expired_keys:
|
||||||
|
del self.intent_cache[key]
|
||||||
|
|
||||||
|
# 如果缓存太大,移除最旧的条目
|
||||||
|
if len(self.intent_cache) > self.cache_max_size:
|
||||||
|
# 按时间戳排序并保留最新的条目
|
||||||
|
sorted_items = sorted(self.intent_cache.items(), key=lambda x: x[1]['timestamp'])
|
||||||
|
for key, _ in sorted_items[:len(sorted_items) - self.cache_max_size]:
|
||||||
|
del self.intent_cache[key]
|
||||||
|
|
||||||
async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str:
|
async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str:
|
||||||
if not self.llm:
|
if not self.llm:
|
||||||
raise ValueError("LLM provider not set")
|
raise ValueError("LLM provider not set")
|
||||||
|
|
||||||
|
# 记录整体开始时间
|
||||||
|
total_start_time = time.time()
|
||||||
|
|
||||||
|
# 打印使用的模型信息
|
||||||
|
model_info = getattr(self.llm, 'model_name', str(self.llm.__class__.__name__))
|
||||||
|
logger.bind(tag=TAG).info(f"使用意图识别模型: {model_info}")
|
||||||
|
|
||||||
|
# 计算缓存键
|
||||||
|
cache_key = hashlib.md5(text.encode()).hexdigest()
|
||||||
|
|
||||||
|
# 检查缓存
|
||||||
|
if cache_key in self.intent_cache:
|
||||||
|
cache_entry = self.intent_cache[cache_key]
|
||||||
|
# 检查缓存是否过期
|
||||||
|
if time.time() - cache_entry['timestamp'] <= self.cache_expiry:
|
||||||
|
cache_time = time.time() - total_start_time
|
||||||
|
logger.bind(tag=TAG).info(f"使用缓存的意图: {cache_key} -> {cache_entry['intent']}, 耗时: {cache_time:.4f}秒")
|
||||||
|
return cache_entry['intent']
|
||||||
|
|
||||||
|
# 清理缓存
|
||||||
|
self.clean_cache()
|
||||||
|
|
||||||
# 构建用户最后一句话的提示
|
# 构建用户最后一句话的提示
|
||||||
msgStr = ""
|
msgStr = ""
|
||||||
|
|
||||||
@@ -94,18 +158,78 @@ class IntentProvider(IntentProviderBase):
|
|||||||
music_file_names = music_config["music_file_names"]
|
music_file_names = music_config["music_file_names"]
|
||||||
prompt_music = f"{self.promot}\n<start>{music_file_names}\n<end>"
|
prompt_music = f"{self.promot}\n<start>{music_file_names}\n<end>"
|
||||||
logger.bind(tag=TAG).debug(f"User prompt: {prompt_music}")
|
logger.bind(tag=TAG).debug(f"User prompt: {prompt_music}")
|
||||||
|
|
||||||
|
# 记录预处理完成时间
|
||||||
|
preprocess_time = time.time() - total_start_time
|
||||||
|
logger.bind(tag=TAG).debug(f"意图识别预处理耗时: {preprocess_time:.4f}秒")
|
||||||
|
|
||||||
# 使用LLM进行意图识别
|
# 使用LLM进行意图识别
|
||||||
|
llm_start_time = time.time()
|
||||||
|
logger.bind(tag=TAG).info(f"开始LLM意图识别调用, 模型: {model_info}")
|
||||||
|
|
||||||
intent = self.llm.response_no_stream(
|
intent = self.llm.response_no_stream(
|
||||||
system_prompt=prompt_music,
|
system_prompt=prompt_music,
|
||||||
user_prompt=user_prompt
|
user_prompt=user_prompt
|
||||||
)
|
)
|
||||||
# 使用正则表达式提取大括号中的内容
|
|
||||||
# 使用正则表达式提取 {} 中的内容
|
# 记录LLM调用完成时间
|
||||||
match = re.search(r'\{.*?\}', intent)
|
llm_time = time.time() - llm_start_time
|
||||||
|
logger.bind(tag=TAG).info(f"LLM意图识别完成, 模型: {model_info}, 调用耗时: {llm_time:.4f}秒")
|
||||||
|
|
||||||
|
# 记录后处理开始时间
|
||||||
|
postprocess_start_time = time.time()
|
||||||
|
|
||||||
|
# 清理和解析响应
|
||||||
|
intent = intent.strip()
|
||||||
|
# 尝试提取JSON部分
|
||||||
|
match = re.search(r'\{.*\}', intent, re.DOTALL)
|
||||||
if match:
|
if match:
|
||||||
result = match.group(0)
|
intent = match.group(0)
|
||||||
intent = result
|
|
||||||
else:
|
# 记录总处理时间
|
||||||
intent = "{intent: '继续聊天'}"
|
total_time = time.time() - total_start_time
|
||||||
logger.bind(tag=TAG).info(f"Detected intent: {intent}")
|
logger.bind(tag=TAG).info(f"【意图识别性能】模型: {model_info}, 总耗时: {total_time:.4f}秒, LLM调用: {llm_time:.4f}秒, 查询: '{text[:20]}...'")
|
||||||
return intent.strip()
|
|
||||||
|
# 尝试解析为JSON
|
||||||
|
try:
|
||||||
|
intent_data = json.loads(intent)
|
||||||
|
# 如果包含function_call,则格式化为适合处理的格式
|
||||||
|
if "function_call" in intent_data:
|
||||||
|
function_data = intent_data["function_call"]
|
||||||
|
function_name = function_data.get("name")
|
||||||
|
function_args = function_data.get("arguments", {})
|
||||||
|
|
||||||
|
# 记录识别到的function call
|
||||||
|
logger.bind(tag=TAG).info(f"识别到function call: {function_name}, 参数: {function_args}")
|
||||||
|
|
||||||
|
# 添加到缓存
|
||||||
|
self.intent_cache[cache_key] = {
|
||||||
|
'intent': intent,
|
||||||
|
'timestamp': time.time()
|
||||||
|
}
|
||||||
|
|
||||||
|
# 后处理时间
|
||||||
|
postprocess_time = time.time() - postprocess_start_time
|
||||||
|
logger.bind(tag=TAG).debug(f"意图后处理耗时: {postprocess_time:.4f}秒")
|
||||||
|
|
||||||
|
# 确保返回完全序列化的JSON字符串
|
||||||
|
return intent
|
||||||
|
else:
|
||||||
|
# 添加到缓存
|
||||||
|
self.intent_cache[cache_key] = {
|
||||||
|
'intent': intent,
|
||||||
|
'timestamp': time.time()
|
||||||
|
}
|
||||||
|
|
||||||
|
# 后处理时间
|
||||||
|
postprocess_time = time.time() - postprocess_start_time
|
||||||
|
logger.bind(tag=TAG).debug(f"意图后处理耗时: {postprocess_time:.4f}秒")
|
||||||
|
|
||||||
|
# 返回普通意图
|
||||||
|
return intent
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
# 后处理时间
|
||||||
|
postprocess_time = time.time() - postprocess_start_time
|
||||||
|
logger.bind(tag=TAG).error(f"无法解析意图JSON: {intent}, 后处理耗时: {postprocess_time:.4f}秒")
|
||||||
|
# 如果解析失败,默认返回继续聊天意图
|
||||||
|
return "{\"intent\": \"继续聊天\"}"
|
||||||
|
|||||||
@@ -0,0 +1,85 @@
|
|||||||
|
from config.logger import setup_logging
|
||||||
|
from openai import OpenAI
|
||||||
|
import json
|
||||||
|
from core.providers.llm.base import LLMProviderBase
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
class LLMProvider(LLMProviderBase):
|
||||||
|
def __init__(self, config):
|
||||||
|
self.model_name = config.get("model_name")
|
||||||
|
self.base_url = config.get("base_url", "http://localhost:9997")
|
||||||
|
# Initialize OpenAI client with Xinference base URL
|
||||||
|
# 如果没有v1,增加v1
|
||||||
|
if not self.base_url.endswith("/v1"):
|
||||||
|
self.base_url = f"{self.base_url}/v1"
|
||||||
|
|
||||||
|
logger.bind(tag=TAG).info(f"Initializing Xinference LLM provider with model: {self.model_name}, base_url: {self.base_url}")
|
||||||
|
|
||||||
|
try:
|
||||||
|
self.client = OpenAI(
|
||||||
|
base_url=self.base_url,
|
||||||
|
api_key="xinference" # Xinference has a similar setup to Ollama where it doesn't need an actual key
|
||||||
|
)
|
||||||
|
logger.bind(tag=TAG).info("Xinference client initialized successfully")
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error initializing Xinference client: {e}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
def response(self, session_id, dialogue):
|
||||||
|
try:
|
||||||
|
logger.bind(tag=TAG).debug(f"Sending request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}")
|
||||||
|
responses = self.client.chat.completions.create(
|
||||||
|
model=self.model_name,
|
||||||
|
messages=dialogue,
|
||||||
|
stream=True
|
||||||
|
)
|
||||||
|
is_active=True
|
||||||
|
for chunk in responses:
|
||||||
|
try:
|
||||||
|
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None
|
||||||
|
content = delta.content if hasattr(delta, 'content') else ''
|
||||||
|
if content:
|
||||||
|
if '<think>' in content:
|
||||||
|
is_active = False
|
||||||
|
content = content.split('<think>')[0]
|
||||||
|
if '</think>' in content:
|
||||||
|
is_active = True
|
||||||
|
content = content.split('</think>')[-1]
|
||||||
|
if is_active:
|
||||||
|
yield content
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error processing chunk: {e}")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error in Xinference response generation: {e}")
|
||||||
|
yield "【Xinference服务响应异常】"
|
||||||
|
|
||||||
|
def response_with_functions(self, session_id, dialogue, functions=None):
|
||||||
|
try:
|
||||||
|
logger.bind(tag=TAG).debug(f"Sending function call request to Xinference with model: {self.model_name}, dialogue length: {len(dialogue)}")
|
||||||
|
if functions:
|
||||||
|
logger.bind(tag=TAG).debug(f"Function calls enabled with: {[f.get('function', {}).get('name') for f in functions]}")
|
||||||
|
|
||||||
|
stream = self.client.chat.completions.create(
|
||||||
|
model=self.model_name,
|
||||||
|
messages=dialogue,
|
||||||
|
stream=True,
|
||||||
|
tools=functions,
|
||||||
|
)
|
||||||
|
|
||||||
|
for chunk in stream:
|
||||||
|
delta = chunk.choices[0].delta
|
||||||
|
content = delta.content
|
||||||
|
tool_calls = delta.tool_calls
|
||||||
|
|
||||||
|
if content:
|
||||||
|
yield content, tool_calls
|
||||||
|
elif tool_calls:
|
||||||
|
yield None, tool_calls
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error in Xinference function call: {e}")
|
||||||
|
yield {"type": "content", "content": f"【Xinference服务响应异常: {str(e)}】"}
|
||||||
@@ -0,0 +1,31 @@
|
|||||||
|
{
|
||||||
|
"des": [
|
||||||
|
"在data目录下创建.mcp_server_settings.json文件,可以选择下面的MCP服务,也可以自行添加新的MCP服务。",
|
||||||
|
"后面不断测试补充好用的mcp服务,欢迎大家一起补充。",
|
||||||
|
"记得删除注释行,des属性仅为说明,不会被解析。",
|
||||||
|
"des和link属性,仅为说明安装方式,方便大家查看原始链接,不是必须项。"
|
||||||
|
],
|
||||||
|
"mcpServers": {
|
||||||
|
"filesystem": {
|
||||||
|
"command": "npx",
|
||||||
|
"args": [
|
||||||
|
"-y",
|
||||||
|
"@modelcontextprotocol/server-filesystem",
|
||||||
|
"/Users/username/Desktop",
|
||||||
|
"/path/to/other/allowed/dir"
|
||||||
|
],
|
||||||
|
"link":"https://github.com/modelcontextprotocol/servers/tree/main/src/filesystem"
|
||||||
|
},
|
||||||
|
"playwright": {
|
||||||
|
"command": "npx",
|
||||||
|
"args": ["-y", "@executeautomation/playwright-mcp-server"],
|
||||||
|
"des" : "run 'npx playwright install' first",
|
||||||
|
"link": "https://github.com/executeautomation/mcp-playwright"
|
||||||
|
},
|
||||||
|
"windows-cli": {
|
||||||
|
"command": "npx",
|
||||||
|
"args": ["-y", "@simonb97/server-win-cli"],
|
||||||
|
"link": "https://github.com/SimonB97/win-cli-mcp-server"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -142,7 +142,8 @@ def get_news(conn, category: str = None, detail: bool = False, lang: str = "zh_C
|
|||||||
detail_content = fetch_news_detail(link)
|
detail_content = fetch_news_detail(link)
|
||||||
|
|
||||||
if not detail_content or detail_content == "无法获取详细内容":
|
if not detail_content or detail_content == "无法获取详细内容":
|
||||||
return ActionResponse(Action.REQLLM, f"抱歉,无法获取《{title}》的详细内容,可能是链接已失效或网站结构发生变化。", None)
|
return ActionResponse(Action.REQLLM,
|
||||||
|
f"抱歉,无法获取《{title}》的详细内容,可能是链接已失效或网站结构发生变化。", None)
|
||||||
|
|
||||||
# 构建详情报告
|
# 构建详情报告
|
||||||
detail_report = (
|
detail_report = (
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ handle_device_function_desc = {
|
|||||||
"description": "动作名称,可选值:get(获取),set(设置),raise(提高),lower(降低)"
|
"description": "动作名称,可选值:get(获取),set(设置),raise(提高),lower(降低)"
|
||||||
},
|
},
|
||||||
"value": {
|
"value": {
|
||||||
"type": "int",
|
"type": "number",
|
||||||
"description": "值大小,可选值:0-100之间的整数"
|
"description": "值大小,可选值:0-100之间的整数"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ class ToolType(Enum):
|
|||||||
CHANGE_SYS_PROMPT = (3, "修改系统提示词,切换角色性格或职责")
|
CHANGE_SYS_PROMPT = (3, "修改系统提示词,切换角色性格或职责")
|
||||||
SYSTEM_CTL = (4, "系统控制,影响正常的对话流程,如退出、播放音乐等,需要传递conn参数")
|
SYSTEM_CTL = (4, "系统控制,影响正常的对话流程,如退出、播放音乐等,需要传递conn参数")
|
||||||
IOT_CTL = (5, "IOT设备控制,需要传递conn参数")
|
IOT_CTL = (5, "IOT设备控制,需要传递conn参数")
|
||||||
|
MCP_CLIENT = (6, "MCP客户端")
|
||||||
|
|
||||||
def __init__(self, code, message):
|
def __init__(self, code, message):
|
||||||
self.code = code
|
self.code = code
|
||||||
|
|||||||
@@ -22,4 +22,6 @@ 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
|
||||||
|
|
||||||
cnlunar==0.2.0
|
cnlunar==0.2.0
|
||||||
Reference in New Issue
Block a user