mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-28 01:53:53 +08:00
本地记忆+意图识别 (#250)
* 增加本地记忆功能,使用llm总结记忆 * update:增加统一非流式输出输出 * 增加意图识别内容,使用llm进行识别 * 初始化记忆模块 * 完善意图识别处理后的流程 * 通过使用function call实现意图识别 * update:优化意图识别的配置 * update:function call最优设置成doubao-pro-32k-functioncall-241028 --------- Co-authored-by: 玄凤科技 <eric230308@gmail.com> Co-authored-by: hrz <1710360675@qq.com>
This commit is contained in:
@@ -65,13 +65,31 @@ CMD_exit:
|
|||||||
|
|
||||||
# 具体处理时选择的模块(The module selected for specific processing)
|
# 具体处理时选择的模块(The module selected for specific processing)
|
||||||
selected_module:
|
selected_module:
|
||||||
ASR: FunASR
|
# 语音活动检测模块,默认使用SileroVAD模型
|
||||||
VAD: SileroVAD
|
VAD: SileroVAD
|
||||||
|
# 语音识别模块,默认使用FunASR本地模型
|
||||||
|
ASR: FunASR
|
||||||
# 将根据配置名称对应的type调用实际的LLM适配器
|
# 将根据配置名称对应的type调用实际的LLM适配器
|
||||||
LLM: ChatGLMLLM
|
LLM: ChatGLMLLM
|
||||||
# TTS将根据配置名称对应的type调用实际的TTS适配器
|
# TTS将根据配置名称对应的type调用实际的TTS适配器
|
||||||
TTS: EdgeTTS
|
TTS: EdgeTTS
|
||||||
Memory: mem0ai
|
# 记忆模块,默认使用本地压缩记忆,如果想使用超长记忆,推荐使用mem0ai
|
||||||
|
Memory: nomem
|
||||||
|
# 意图识别模块,默认不开启。开启后,可以播放音乐、控制音量、识别退出指令
|
||||||
|
# 如果意图识别设置成 function_call,建议使用:DoubaoLLM,推荐把的model_name设置为:doubao-pro-32k-functioncall-241028
|
||||||
|
Intent: nointent
|
||||||
|
|
||||||
|
# 意图识别,是用于理解用户意图的模块,例如:播放音乐
|
||||||
|
Intent:
|
||||||
|
# 不使用意图识别
|
||||||
|
nointent:
|
||||||
|
type: nointent
|
||||||
|
# 意图识别使用大模型,优点:通用性强,缺点:增加串行前置意图识别模块,会增加处理时间
|
||||||
|
intent_llm:
|
||||||
|
type: intent_llm
|
||||||
|
# 意图识别使用function_call,缺点:需要所选择的LLM支持function_call,优点:按需调用工具、速度快
|
||||||
|
function_call:
|
||||||
|
type: function_call
|
||||||
|
|
||||||
Memory:
|
Memory:
|
||||||
mem0ai:
|
mem0ai:
|
||||||
@@ -79,6 +97,12 @@ Memory:
|
|||||||
# https://app.mem0.ai/dashboard/api-keys
|
# https://app.mem0.ai/dashboard/api-keys
|
||||||
# 每月有1000次免费调用
|
# 每月有1000次免费调用
|
||||||
api_key: 你的mem0ai api key
|
api_key: 你的mem0ai api key
|
||||||
|
nomem:
|
||||||
|
# 不想使用记忆功能,可以使用nomem
|
||||||
|
type: nomem
|
||||||
|
mem_local_short:
|
||||||
|
# 本地记忆功能,通过selected_module的llm总结,数据保存在本地,不会上传到服务器
|
||||||
|
type: mem_local_short
|
||||||
|
|
||||||
ASR:
|
ASR:
|
||||||
FunASR:
|
FunASR:
|
||||||
@@ -106,6 +130,13 @@ LLM:
|
|||||||
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
|
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
|
||||||
model_name: qwen-turbo
|
model_name: qwen-turbo
|
||||||
api_key: 你的deepseek web key
|
api_key: 你的deepseek web key
|
||||||
|
DoubaoLLM:
|
||||||
|
# 定义LLM API类型
|
||||||
|
type: openai
|
||||||
|
# 可在这里开通 https://console.volcengine.com/ark/region:ark+cn-beijing/model/detail?Id=doubao-pro-32k&projectName=undefined
|
||||||
|
base_url: https://ark.cn-beijing.volces.com/api/v3
|
||||||
|
model_name: doubao-pro-32k-functioncall-241028
|
||||||
|
api_key: 你的doubao web key
|
||||||
DeepSeekLLM:
|
DeepSeekLLM:
|
||||||
# 定义LLM API类型
|
# 定义LLM API类型
|
||||||
type: openai
|
type: openai
|
||||||
@@ -383,18 +414,6 @@ module_test:
|
|||||||
|
|
||||||
# 本地音乐播放配置
|
# 本地音乐播放配置
|
||||||
music:
|
music:
|
||||||
music_commands:
|
|
||||||
- "来一首歌"
|
|
||||||
- "唱一首歌"
|
|
||||||
- "播放音乐"
|
|
||||||
- "来点音乐"
|
|
||||||
- "背景音乐"
|
|
||||||
- "放首歌"
|
|
||||||
- "播放歌曲"
|
|
||||||
- "来点背景音乐"
|
|
||||||
- "我想听歌"
|
|
||||||
- "我要听歌"
|
|
||||||
- "放点音乐"
|
|
||||||
music_dir: "./music" # 音乐文件存放路径,将从该目录及子目录下搜索音乐文件
|
music_dir: "./music" # 音乐文件存放路径,将从该目录及子目录下搜索音乐文件
|
||||||
music_ext: # 音乐文件类型,p3格式效率最高
|
music_ext: # 音乐文件类型,p3格式效率最高
|
||||||
- ".mp3"
|
- ".mp3"
|
||||||
|
|||||||
@@ -0,0 +1,36 @@
|
|||||||
|
FunctionCallConfig = [
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "handle_exit_intent",
|
||||||
|
"description": "当用户想结束对话或需要退出系统时调用",
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"say_goodbye": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "和用户友好结束对话的告别语"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"required": []
|
||||||
|
}
|
||||||
|
}
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"type": "function",
|
||||||
|
"function": {
|
||||||
|
"name": "play_music",
|
||||||
|
"description": "唱歌、听歌、播放音乐方法。比如用户说播放音乐,参数为:random,比如用户说播放两只老虎,参数为:两只老虎",
|
||||||
|
"parameters": {
|
||||||
|
"type": "object",
|
||||||
|
"properties": {
|
||||||
|
"song_name": {
|
||||||
|
"type": "string",
|
||||||
|
"description": "歌曲名称,如果没有指定具体歌名则为'random'"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"required": ["song_name"]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
@@ -15,6 +15,7 @@ from core.utils.util import get_string_no_punctuation_or_emoji
|
|||||||
from concurrent.futures import ThreadPoolExecutor, TimeoutError
|
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.intentHandler import Action, get_functions, handle_llm_function_call
|
||||||
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
|
||||||
@@ -27,7 +28,7 @@ class TTSException(RuntimeError):
|
|||||||
|
|
||||||
|
|
||||||
class ConnectionHandler:
|
class ConnectionHandler:
|
||||||
def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _music, _memory):
|
def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _music, _memory, _intent):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.logger = setup_logging()
|
self.logger = setup_logging()
|
||||||
self.auth = AuthMiddleware(config)
|
self.auth = AuthMiddleware(config)
|
||||||
@@ -55,6 +56,7 @@ class ConnectionHandler:
|
|||||||
self.llm = _llm
|
self.llm = _llm
|
||||||
self.tts = _tts
|
self.tts = _tts
|
||||||
self.memory = _memory
|
self.memory = _memory
|
||||||
|
self.intent = _intent
|
||||||
|
|
||||||
# vad相关变量
|
# vad相关变量
|
||||||
self.client_audio_buffer = bytes()
|
self.client_audio_buffer = bytes()
|
||||||
@@ -88,6 +90,12 @@ class ConnectionHandler:
|
|||||||
self.auth_code_gen = AuthCodeGenerator.get_instance()
|
self.auth_code_gen = AuthCodeGenerator.get_instance()
|
||||||
self.is_device_verified = False # 添加设备验证状态标志
|
self.is_device_verified = False # 添加设备验证状态标志
|
||||||
self.music_handler = _music
|
self.music_handler = _music
|
||||||
|
self.close_after_chat = False # 是否在聊天结束后关闭连接
|
||||||
|
self.use_function_call_mode = False
|
||||||
|
if self.config["selected_module"]["Intent"] == 'function_call':
|
||||||
|
self.use_function_call_mode = True
|
||||||
|
|
||||||
|
self.logger.bind(tag=TAG).info(f"use_function_call_mode:{self.use_function_call_mode}")
|
||||||
|
|
||||||
async def handle_connection(self, ws):
|
async def handle_connection(self, ws):
|
||||||
try:
|
try:
|
||||||
@@ -101,7 +109,8 @@ class ConnectionHandler:
|
|||||||
await self.auth.authenticate(self.headers)
|
await self.auth.authenticate(self.headers)
|
||||||
|
|
||||||
device_id = self.headers.get("device-id", None)
|
device_id = self.headers.get("device-id", None)
|
||||||
self.memory.set_role_id(device_id)
|
self.memory.init_memory(device_id, self.llm)
|
||||||
|
self.intent.set_llm(self.llm)
|
||||||
|
|
||||||
# Load private configuration if device_id is provided
|
# Load private configuration if device_id is provided
|
||||||
bUsePrivateConfig = self.config.get("use_private_config", False)
|
bUsePrivateConfig = self.config.get("use_private_config", False)
|
||||||
@@ -204,7 +213,6 @@ class ConnectionHandler:
|
|||||||
return False
|
return False
|
||||||
return not self.is_device_verified
|
return not self.is_device_verified
|
||||||
|
|
||||||
|
|
||||||
def chat(self, query):
|
def chat(self, query):
|
||||||
if self.isNeedAuth():
|
if self.isNeedAuth():
|
||||||
self.llm_finish_task = True
|
self.llm_finish_task = True
|
||||||
@@ -213,6 +221,7 @@ class ConnectionHandler:
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
self.dialogue.put(Message(role="user", content=query))
|
self.dialogue.put(Message(role="user", content=query))
|
||||||
|
|
||||||
response_message = []
|
response_message = []
|
||||||
processed_chars = 0 # 跟踪已处理的字符位置
|
processed_chars = 0 # 跟踪已处理的字符位置
|
||||||
try:
|
try:
|
||||||
@@ -221,7 +230,7 @@ class ConnectionHandler:
|
|||||||
future = asyncio.run_coroutine_threadsafe(self.memory.query_memory(query), self.loop)
|
future = asyncio.run_coroutine_threadsafe(self.memory.query_memory(query), self.loop)
|
||||||
memory_str = future.result()
|
memory_str = future.result()
|
||||||
|
|
||||||
self.logger.bind(tag=TAG).info(f"记忆内容: {memory_str}")
|
self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}")
|
||||||
llm_responses = self.llm.response(
|
llm_responses = self.llm.response(
|
||||||
self.session_id,
|
self.session_id,
|
||||||
self.dialogue.get_llm_dialogue_with_memory(memory_str)
|
self.dialogue.get_llm_dialogue_with_memory(memory_str)
|
||||||
@@ -245,7 +254,7 @@ class ConnectionHandler:
|
|||||||
current_text = full_text[processed_chars:] # 从未处理的位置开始
|
current_text = full_text[processed_chars:] # 从未处理的位置开始
|
||||||
|
|
||||||
# 查找最后一个有效标点
|
# 查找最后一个有效标点
|
||||||
punctuations = ("。", "?", "!", "?", "!", ";", ";", ":", ":")
|
punctuations = ("。", "?", "!", ";", ":")
|
||||||
last_punct_pos = -1
|
last_punct_pos = -1
|
||||||
for punct in punctuations:
|
for punct in punctuations:
|
||||||
pos = current_text.rfind(punct)
|
pos = current_text.rfind(punct)
|
||||||
@@ -282,6 +291,119 @@ class ConnectionHandler:
|
|||||||
self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False))
|
self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False))
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
def chat_with_function_calling(self, query):
|
||||||
|
self.logger.bind(tag=TAG).debug(f"Chat with function calling start: {query}")
|
||||||
|
"""Chat with function calling for intent detection using streaming"""
|
||||||
|
if self.isNeedAuth():
|
||||||
|
self.llm_finish_task = True
|
||||||
|
future = asyncio.run_coroutine_threadsafe(self._check_and_broadcast_auth_code(), self.loop)
|
||||||
|
future.result()
|
||||||
|
return True
|
||||||
|
|
||||||
|
self.dialogue.put(Message(role="user", content=query))
|
||||||
|
|
||||||
|
# Define intent functions
|
||||||
|
functions = get_functions()
|
||||||
|
|
||||||
|
response_message = []
|
||||||
|
processed_chars = 0 # 跟踪已处理的字符位置
|
||||||
|
function_call_data = None # 存储function call数据
|
||||||
|
|
||||||
|
try:
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
|
# 使用带记忆的对话
|
||||||
|
future = asyncio.run_coroutine_threadsafe(self.memory.query_memory(query), self.loop)
|
||||||
|
memory_str = future.result()
|
||||||
|
|
||||||
|
# self.logger.bind(tag=TAG).info(f"记忆内容: {memory_str}")
|
||||||
|
|
||||||
|
# 使用支持functions的streaming接口
|
||||||
|
llm_responses = self.llm.response_with_functions(
|
||||||
|
self.session_id,
|
||||||
|
self.dialogue.get_llm_dialogue_with_memory(memory_str),
|
||||||
|
functions=functions
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"LLM 处理出错 {query}: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
self.llm_finish_task = False
|
||||||
|
text_index = 0
|
||||||
|
|
||||||
|
# 处理流式响应
|
||||||
|
for response in llm_responses:
|
||||||
|
if response["type"] == "content":
|
||||||
|
content = response["content"]
|
||||||
|
response_message.append(content)
|
||||||
|
|
||||||
|
if self.client_abort:
|
||||||
|
break
|
||||||
|
|
||||||
|
end_time = time.time()
|
||||||
|
self.logger.bind(tag=TAG).debug(f"大模型返回时间: {end_time - start_time} 秒, 生成token={content}")
|
||||||
|
|
||||||
|
# 处理文本分段和TTS逻辑
|
||||||
|
# 合并当前全部文本并处理未分割部分
|
||||||
|
full_text = "".join(response_message)
|
||||||
|
current_text = full_text[processed_chars:] # 从未处理的位置开始
|
||||||
|
|
||||||
|
# 查找最后一个有效标点
|
||||||
|
punctuations = ("。", "?", "!", ";", ":")
|
||||||
|
last_punct_pos = -1
|
||||||
|
for punct in punctuations:
|
||||||
|
pos = current_text.rfind(punct)
|
||||||
|
if pos > last_punct_pos:
|
||||||
|
last_punct_pos = pos
|
||||||
|
|
||||||
|
# 找到分割点则处理
|
||||||
|
if last_punct_pos != -1:
|
||||||
|
segment_text_raw = current_text[:last_punct_pos + 1]
|
||||||
|
segment_text = get_string_no_punctuation_or_emoji(segment_text_raw)
|
||||||
|
if segment_text:
|
||||||
|
text_index += 1
|
||||||
|
self.recode_first_last_text(segment_text, text_index)
|
||||||
|
future = self.executor.submit(self.speak_and_play, segment_text, text_index)
|
||||||
|
self.tts_queue.put(future)
|
||||||
|
processed_chars += len(segment_text_raw) # 更新已处理字符位置
|
||||||
|
|
||||||
|
elif response["type"] == "function_call":
|
||||||
|
# Extract function call data
|
||||||
|
function_call_data = {
|
||||||
|
"name": response["function_call"]["function"]["name"],
|
||||||
|
"arguments": response["function_call"]["function"]["arguments"]
|
||||||
|
}
|
||||||
|
self.logger.bind(tag=TAG).info(f"Function call detected: {function_call_data}")
|
||||||
|
|
||||||
|
# 处理最后剩余的文本
|
||||||
|
full_text = "".join(response_message)
|
||||||
|
remaining_text = full_text[processed_chars:]
|
||||||
|
if remaining_text:
|
||||||
|
segment_text = get_string_no_punctuation_or_emoji(remaining_text)
|
||||||
|
if segment_text:
|
||||||
|
text_index += 1
|
||||||
|
self.recode_first_last_text(segment_text, text_index)
|
||||||
|
future = self.executor.submit(self.speak_and_play, segment_text, text_index)
|
||||||
|
self.tts_queue.put(future)
|
||||||
|
|
||||||
|
# 存储对话内容
|
||||||
|
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
|
||||||
|
|
||||||
|
# 处理function call
|
||||||
|
if function_call_data:
|
||||||
|
result = handle_llm_function_call(self, function_call_data)
|
||||||
|
if result.action == Action.RESPONSE:
|
||||||
|
text = result.response
|
||||||
|
text_index += 1
|
||||||
|
self.recode_first_last_text(text, text_index)
|
||||||
|
future = self.executor.submit(self.speak_and_play, text, text_index)
|
||||||
|
self.tts_queue.put(future)
|
||||||
|
|
||||||
|
self.llm_finish_task = True
|
||||||
|
self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False))
|
||||||
|
|
||||||
|
return True
|
||||||
|
|
||||||
def _tts_priority_thread(self):
|
def _tts_priority_thread(self):
|
||||||
while not self.stop_event.is_set():
|
while not self.stop_event.is_set():
|
||||||
text = None
|
text = None
|
||||||
@@ -372,3 +494,14 @@ class ConnectionHandler:
|
|||||||
self.client_have_voice_last_time = 0
|
self.client_have_voice_last_time = 0
|
||||||
self.client_voice_stop = False
|
self.client_voice_stop = False
|
||||||
self.logger.bind(tag=TAG).debug("VAD states reset.")
|
self.logger.bind(tag=TAG).debug("VAD states reset.")
|
||||||
|
|
||||||
|
def chat_and_close(self, text):
|
||||||
|
"""Chat with the user and then close the connection"""
|
||||||
|
try:
|
||||||
|
# Use the existing chat method
|
||||||
|
self.chat(text)
|
||||||
|
|
||||||
|
# After chat is complete, close the connection
|
||||||
|
self.close_after_chat = True
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"Chat and close error: {str(e)}")
|
||||||
|
|||||||
@@ -0,0 +1,170 @@
|
|||||||
|
from config.logger import setup_logging
|
||||||
|
import json
|
||||||
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
|
from core.utils.dialogue import Message
|
||||||
|
from config.functionCallConfig import FunctionCallConfig
|
||||||
|
import asyncio
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
class Action(Enum):
|
||||||
|
NOTFOUND = (0, "没有找到函数")
|
||||||
|
NONE = (1, "啥也不干")
|
||||||
|
RESPONSE = (2, "直接回复")
|
||||||
|
REQLLM = (3, "调用函数后再请求llm生成回复")
|
||||||
|
|
||||||
|
def __init__(self, code, message):
|
||||||
|
self.code = code
|
||||||
|
self.message = message
|
||||||
|
|
||||||
|
|
||||||
|
class ActionResponse:
|
||||||
|
def __init__(self, action: Action, result, response):
|
||||||
|
self.action = action # 动作类型
|
||||||
|
self.result = result # 动作产生的结果
|
||||||
|
self.response = response # 直接回复的内容
|
||||||
|
|
||||||
|
|
||||||
|
def get_functions():
|
||||||
|
"""获取功能调用配置"""
|
||||||
|
return FunctionCallConfig
|
||||||
|
|
||||||
|
|
||||||
|
def handle_llm_function_call(conn, function_call_data):
|
||||||
|
try:
|
||||||
|
function_name = function_call_data["name"]
|
||||||
|
|
||||||
|
if function_name == "handle_exit_intent":
|
||||||
|
# 处理退出意图
|
||||||
|
try:
|
||||||
|
say_goodbye = json.loads(function_call_data["arguments"]).get("say_goodbye", "再见")
|
||||||
|
conn.close_after_chat = True
|
||||||
|
logger.bind(tag=TAG).info(f"退出意图已处理:{say_goodbye}")
|
||||||
|
return ActionResponse(action=Action.RESPONSE, result="退出意图已处理", response=say_goodbye)
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"处理退出意图错误: {e}")
|
||||||
|
|
||||||
|
elif function_name == "play_music":
|
||||||
|
# 处理音乐播放意图
|
||||||
|
try:
|
||||||
|
song_name = "random"
|
||||||
|
arguments = function_call_data["arguments"]
|
||||||
|
if arguments is not None and len(arguments) > 0:
|
||||||
|
args = json.loads(arguments)
|
||||||
|
song_name = args.get("song_name", "random")
|
||||||
|
music_intent = f"播放音乐 {song_name}" if song_name != "random" else "随机播放音乐"
|
||||||
|
|
||||||
|
# 执行音乐播放命令
|
||||||
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
|
conn.music_handler.handle_music_command(conn, music_intent),
|
||||||
|
conn.loop
|
||||||
|
)
|
||||||
|
future.result()
|
||||||
|
return ActionResponse(action=Action.RESPONSE, result="退出意图已处理", response="还想听什么歌?")
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"处理音乐意图错误: {e}")
|
||||||
|
else:
|
||||||
|
return ActionResponse(action=Action.NOTFOUND, result="没有找到对应的函数", response="")
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"处理function call错误: {e}")
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def handle_user_intent(conn, text):
|
||||||
|
"""
|
||||||
|
Handle user intent before starting chat
|
||||||
|
|
||||||
|
Args:
|
||||||
|
conn: Connection object
|
||||||
|
text: User's text input
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: True if intent was handled, False if should proceed to chat
|
||||||
|
"""
|
||||||
|
# 检查是否有明确的退出命令
|
||||||
|
if await check_direct_exit(conn, text):
|
||||||
|
return True
|
||||||
|
|
||||||
|
if conn.use_function_call_mode:
|
||||||
|
# 使用支持function calling的聊天方法,不再进行意图分析
|
||||||
|
return False
|
||||||
|
|
||||||
|
logger.bind(tag=TAG).info(f"分析用户意图: {text}")
|
||||||
|
|
||||||
|
# 使用LLM进行意图分析
|
||||||
|
intent = await analyze_intent_with_llm(conn, text)
|
||||||
|
|
||||||
|
if not intent:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 处理各种意图
|
||||||
|
return await process_intent_result(conn, intent, text)
|
||||||
|
|
||||||
|
|
||||||
|
async def check_direct_exit(conn, text):
|
||||||
|
"""检查是否有明确的退出命令"""
|
||||||
|
cmd_exit = conn.cmd_exit
|
||||||
|
for cmd in cmd_exit:
|
||||||
|
if text == cmd:
|
||||||
|
logger.bind(tag=TAG).info(f"识别到明确的退出命令: {text}")
|
||||||
|
await conn.close()
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
async def analyze_intent_with_llm(conn, text):
|
||||||
|
"""使用LLM分析用户意图"""
|
||||||
|
if not hasattr(conn, 'intent') or not conn.intent:
|
||||||
|
logger.bind(tag=TAG).warning("意图识别服务未初始化")
|
||||||
|
return None
|
||||||
|
|
||||||
|
# 创建对话历史记录
|
||||||
|
dialogue = conn.dialogue
|
||||||
|
dialogue.put(Message(role="user", content=text))
|
||||||
|
|
||||||
|
try:
|
||||||
|
intent_result = await conn.intent.detect_intent(dialogue.dialogue)
|
||||||
|
logger.bind(tag=TAG).info(f"意图识别结果: {intent_result}")
|
||||||
|
|
||||||
|
# 尝试解析JSON结果
|
||||||
|
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:
|
||||||
|
logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}")
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
async def process_intent_result(conn, intent, original_text):
|
||||||
|
"""处理意图识别结果"""
|
||||||
|
# 处理退出意图
|
||||||
|
if "结束聊天" in intent:
|
||||||
|
logger.bind(tag=TAG).info(f"识别到退出意图: {intent}")
|
||||||
|
|
||||||
|
# 如果正在播放音乐,可以关了 TODO
|
||||||
|
|
||||||
|
# 如果是明确的离别意图,发送告别语并关闭连接
|
||||||
|
await send_stt_message(conn, original_text)
|
||||||
|
conn.executor.submit(conn.chat_and_close, original_text)
|
||||||
|
return True
|
||||||
|
|
||||||
|
# 处理播放音乐意图
|
||||||
|
if "播放音乐" in intent:
|
||||||
|
logger.bind(tag=TAG).info(f"识别到音乐播放意图: {intent}")
|
||||||
|
await conn.music_handler.handle_music_command(conn, intent)
|
||||||
|
return True
|
||||||
|
|
||||||
|
# 其他意图处理可以在这里扩展
|
||||||
|
|
||||||
|
# 默认返回False,表示继续常规聊天流程
|
||||||
|
return False
|
||||||
@@ -15,7 +15,7 @@ logger = setup_logging()
|
|||||||
|
|
||||||
def _extract_song_name(text):
|
def _extract_song_name(text):
|
||||||
"""从用户输入中提取歌名"""
|
"""从用户输入中提取歌名"""
|
||||||
for keyword in ["听", "播放", "放", "唱"]:
|
for keyword in ["播放音乐"]:
|
||||||
if keyword in text:
|
if keyword in text:
|
||||||
parts = text.split(keyword)
|
parts = text.split(keyword)
|
||||||
if len(parts) > 1:
|
if len(parts) > 1:
|
||||||
@@ -36,6 +36,7 @@ def _find_best_match(potential_song, music_files):
|
|||||||
best_match = music_file
|
best_match = music_file
|
||||||
return best_match
|
return best_match
|
||||||
|
|
||||||
|
|
||||||
class MusicManager:
|
class MusicManager:
|
||||||
def __init__(self, music_dir, music_ext):
|
def __init__(self, music_dir, music_ext):
|
||||||
self.music_dir = Path(music_dir)
|
self.music_dir = Path(music_dir)
|
||||||
@@ -55,23 +56,20 @@ class MusicManager:
|
|||||||
music_files.append(str(file.relative_to(self.music_dir)))
|
music_files.append(str(file.relative_to(self.music_dir)))
|
||||||
return music_files
|
return music_files
|
||||||
|
|
||||||
|
|
||||||
class MusicHandler:
|
class MusicHandler:
|
||||||
def __init__(self, config):
|
def __init__(self, config):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.music_related_keywords = []
|
|
||||||
|
|
||||||
if "music" in self.config:
|
if "music" in self.config:
|
||||||
self.music_config = self.config["music"]
|
self.music_config = self.config["music"]
|
||||||
self.music_dir = os.path.abspath(
|
self.music_dir = os.path.abspath(
|
||||||
self.music_config.get("music_dir", "./music") # 默认路径修改
|
self.music_config.get("music_dir", "./music") # 默认路径修改
|
||||||
)
|
)
|
||||||
self.music_related_keywords = self.music_config.get("music_commands", [])
|
|
||||||
self.music_ext = self.music_config.get("music_ext", (".mp3", ".wav", ".p3"))
|
self.music_ext = self.music_config.get("music_ext", (".mp3", ".wav", ".p3"))
|
||||||
self.refresh_time = self.music_config.get("refresh_time", 60)
|
self.refresh_time = self.music_config.get("refresh_time", 60)
|
||||||
else:
|
else:
|
||||||
self.music_dir = os.path.abspath("./music")
|
self.music_dir = os.path.abspath("./music")
|
||||||
self.music_related_keywords = ["来一首歌", "唱一首歌", "播放音乐", "来点音乐", "背景音乐", "放首歌",
|
|
||||||
"播放歌曲", "来点背景音乐", "我想听歌", "我要听歌", "放点音乐"]
|
|
||||||
self.music_ext = (".mp3", ".wav", ".p3")
|
self.music_ext = (".mp3", ".wav", ".p3")
|
||||||
self.refresh_time = 60
|
self.refresh_time = 60
|
||||||
|
|
||||||
@@ -100,13 +98,9 @@ class MusicHandler:
|
|||||||
logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}")
|
logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}")
|
||||||
await self.play_local_music(conn, specific_file=best_match)
|
await self.play_local_music(conn, specific_file=best_match)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# 检查是否是通用播放音乐命令
|
# 检查是否是通用播放音乐命令
|
||||||
if any(cmd in clean_text for cmd in self.music_related_keywords):
|
await self.play_local_music(conn)
|
||||||
await self.play_local_music(conn)
|
return True
|
||||||
return True
|
|
||||||
|
|
||||||
return False
|
|
||||||
|
|
||||||
async def play_local_music(self, conn, specific_file=None):
|
async def play_local_music(self, conn, specific_file=None):
|
||||||
"""播放本地音乐文件"""
|
"""播放本地音乐文件"""
|
||||||
@@ -117,26 +111,18 @@ class MusicHandler:
|
|||||||
|
|
||||||
# 确保路径正确性
|
# 确保路径正确性
|
||||||
if specific_file:
|
if specific_file:
|
||||||
music_path = os.path.join(self.music_dir, specific_file)
|
|
||||||
if not os.path.exists(music_path):
|
|
||||||
logger.bind(tag=TAG).error(f"指定的音乐文件不存在: {music_path}")
|
|
||||||
return
|
|
||||||
selected_music = specific_file
|
selected_music = specific_file
|
||||||
|
music_path = os.path.join(self.music_dir, specific_file)
|
||||||
else:
|
else:
|
||||||
if time.time() - self.scan_time > self.refresh_time:
|
|
||||||
# 刷新音乐文件列表
|
|
||||||
self.music_files = MusicManager(self.music_dir, self.music_ext).get_music_files()
|
|
||||||
self.scan_time = time.time()
|
|
||||||
logger.bind(tag=TAG).debug(f"刷新的音乐文件列表: {self.music_files}")
|
|
||||||
|
|
||||||
if not self.music_files:
|
if not self.music_files:
|
||||||
logger.bind(tag=TAG).error("未找到MP3音乐文件")
|
logger.bind(tag=TAG).error("未找到MP3音乐文件")
|
||||||
return
|
return
|
||||||
selected_music = random.choice(self.music_files)
|
selected_music = random.choice(self.music_files)
|
||||||
music_path = os.path.join(self.music_dir, selected_music)
|
music_path = os.path.join(self.music_dir, selected_music)
|
||||||
if not os.path.exists(music_path):
|
|
||||||
logger.bind(tag=TAG).error(f"选定的音乐文件不存在: {music_path}")
|
if not os.path.exists(music_path):
|
||||||
return
|
logger.bind(tag=TAG).error(f"选定的音乐文件不存在: {music_path}")
|
||||||
|
return
|
||||||
text = f"正在播放{selected_music}"
|
text = f"正在播放{selected_music}"
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
conn.tts_first_text_index = 0
|
conn.tts_first_text_index = 0
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from config.logger import setup_logging
|
|||||||
import time
|
import time
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import send_stt_message
|
||||||
|
from core.handle.intentHandler import handle_user_intent
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -33,13 +34,7 @@ async def handleAudioMessage(conn, audio):
|
|||||||
else:
|
else:
|
||||||
text, file_path = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
|
text, file_path = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
|
||||||
logger.bind(tag=TAG).info(f"识别文本: {text}")
|
logger.bind(tag=TAG).info(f"识别文本: {text}")
|
||||||
text_len, text_without_punctuation = remove_punctuation_and_length(text)
|
text_len, _ = remove_punctuation_and_length(text)
|
||||||
if await conn.music_handler.handle_music_command(conn, text_without_punctuation):
|
|
||||||
conn.asr_server_receive = True
|
|
||||||
conn.asr_audio.clear()
|
|
||||||
return
|
|
||||||
if text_len <= conn.max_cmd_length and await handleCMDMessage(conn, text_without_punctuation):
|
|
||||||
return
|
|
||||||
if text_len > 0:
|
if text_len > 0:
|
||||||
await startToChat(conn, text)
|
await startToChat(conn, text)
|
||||||
else:
|
else:
|
||||||
@@ -48,20 +43,22 @@ async def handleAudioMessage(conn, audio):
|
|||||||
conn.reset_vad_states()
|
conn.reset_vad_states()
|
||||||
|
|
||||||
|
|
||||||
async def handleCMDMessage(conn, text):
|
|
||||||
cmd_exit = conn.cmd_exit
|
|
||||||
for cmd in cmd_exit:
|
|
||||||
if text == cmd:
|
|
||||||
logger.bind(tag=TAG).info("识别到明确的退出命令".format(text))
|
|
||||||
await conn.close()
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
async def startToChat(conn, text):
|
async def startToChat(conn, text):
|
||||||
# 异步发送 stt 信息
|
# 首先进行意图分析
|
||||||
|
intent_handled = await handle_user_intent(conn, text)
|
||||||
|
|
||||||
|
if intent_handled:
|
||||||
|
# 如果意图已被处理,不再进行聊天
|
||||||
|
conn.asr_server_receive = True
|
||||||
|
return
|
||||||
|
|
||||||
|
# 意图未被处理,继续常规聊天流程
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
conn.executor.submit(conn.chat, text)
|
if conn.use_function_call_mode:
|
||||||
|
# 使用支持function calling的聊天方法
|
||||||
|
conn.executor.submit(conn.chat_with_function_calling, text)
|
||||||
|
else:
|
||||||
|
conn.executor.submit(conn.chat, text)
|
||||||
|
|
||||||
|
|
||||||
async def no_voice_close_connect(conn):
|
async def no_voice_close_connect(conn):
|
||||||
|
|||||||
@@ -7,14 +7,6 @@ from core.utils.util import remove_punctuation_and_length, get_string_no_punctua
|
|||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
async def isLLMWantToFinish(last_text):
|
|
||||||
_, last_text_without_punctuation = remove_punctuation_and_length(last_text)
|
|
||||||
if "再见" in last_text_without_punctuation or "拜拜" in last_text_without_punctuation:
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
async def sendAudioMessage(conn, audios, text, text_index=0):
|
async def sendAudioMessage(conn, audios, text, text_index=0):
|
||||||
# 发送句子开始消息
|
# 发送句子开始消息
|
||||||
if text_index == conn.tts_first_text_index:
|
if text_index == conn.tts_first_text_index:
|
||||||
@@ -46,10 +38,9 @@ async def sendAudioMessage(conn, audios, text, text_index=0):
|
|||||||
# 发送结束消息(如果是最后一个文本)
|
# 发送结束消息(如果是最后一个文本)
|
||||||
if conn.llm_finish_task and text_index == conn.tts_last_text_index:
|
if conn.llm_finish_task and text_index == conn.tts_last_text_index:
|
||||||
await send_tts_message(conn, 'stop', None)
|
await send_tts_message(conn, 'stop', None)
|
||||||
if await isLLMWantToFinish(text):
|
if conn.close_after_chat:
|
||||||
await conn.close()
|
await conn.close()
|
||||||
|
|
||||||
|
|
||||||
async def send_tts_message(conn, state, text=None):
|
async def send_tts_message(conn, state, text=None):
|
||||||
"""发送 TTS 状态消息"""
|
"""发送 TTS 状态消息"""
|
||||||
message = {
|
message = {
|
||||||
|
|||||||
@@ -0,0 +1,33 @@
|
|||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from typing import List, Dict
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
class IntentProviderBase(ABC):
|
||||||
|
def __init__(self, config):
|
||||||
|
self.config = config
|
||||||
|
self.intent_options = config.get("intent_options", {
|
||||||
|
"continue_chat": "继续聊天",
|
||||||
|
"end_chat": "结束聊天",
|
||||||
|
"play_music": "播放音乐"
|
||||||
|
})
|
||||||
|
|
||||||
|
def set_llm(self, llm):
|
||||||
|
self.llm = llm
|
||||||
|
logger.bind(tag=TAG).debug("Set LLM for intent provider")
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
async def detect_intent(self, dialogue_history: List[Dict]) -> str:
|
||||||
|
"""
|
||||||
|
检测用户最后一句话的意图
|
||||||
|
Args:
|
||||||
|
dialogue_history: 对话历史记录列表,每条记录包含role和content
|
||||||
|
Returns:
|
||||||
|
返回识别出的意图,格式为:
|
||||||
|
- "继续聊天"
|
||||||
|
- "结束聊天"
|
||||||
|
- "播放音乐 歌名" 或 "随机播放音乐"
|
||||||
|
"""
|
||||||
|
pass
|
||||||
@@ -0,0 +1,61 @@
|
|||||||
|
from typing import List, Dict
|
||||||
|
from ..base import IntentProviderBase
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
class IntentProvider(IntentProviderBase):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__(config)
|
||||||
|
self.llm = None
|
||||||
|
self.promot = self.get_intent_system_prompt()
|
||||||
|
|
||||||
|
def get_intent_system_prompt(self) -> str:
|
||||||
|
"""
|
||||||
|
根据配置的意图选项动态生成系统提示词
|
||||||
|
Returns:
|
||||||
|
格式化后的系统提示词
|
||||||
|
"""
|
||||||
|
intent_list = []
|
||||||
|
for key, value in self.intent_options.items():
|
||||||
|
if key == "play_music":
|
||||||
|
intent_list.append(f"{value} [歌名]")
|
||||||
|
else:
|
||||||
|
intent_list.append(value)
|
||||||
|
|
||||||
|
prompt = (
|
||||||
|
"你是一个意图识别助手。你需要根据和用户的对话记录,重点分析用户的最后一句话,判断用户意图属于以下哪一类:\n"
|
||||||
|
f"{', '.join(intent_list)}\n"
|
||||||
|
"如果是唱歌、听歌、播放音乐,请指定歌名,格式为'播放音乐 [识别出的歌名]'。\n"
|
||||||
|
"如果听不出具体歌名,可以返回'随机播放音乐'。\n"
|
||||||
|
"只需要返回意图结果的json,不要解释。"
|
||||||
|
"返回格式如下:\n"
|
||||||
|
"{intent: '用户意图'}"
|
||||||
|
)
|
||||||
|
return prompt
|
||||||
|
|
||||||
|
async def detect_intent(self, dialogue_history: List[Dict]) -> str:
|
||||||
|
if not self.llm:
|
||||||
|
raise ValueError("LLM provider not set")
|
||||||
|
|
||||||
|
# 构建用户最后一句话的提示
|
||||||
|
msgStr = ""
|
||||||
|
for msg in dialogue_history:
|
||||||
|
if msg.role == "user":
|
||||||
|
msgStr += f"User: {msg.content}\n"
|
||||||
|
elif msg.role== "assistant":
|
||||||
|
msgStr += f"Assistant: {msg.content}\n"
|
||||||
|
|
||||||
|
user_prompt = f"请分析用户的意图:\n{msgStr}"
|
||||||
|
|
||||||
|
# 使用LLM进行意图识别
|
||||||
|
intent = self.llm.response_no_stream(
|
||||||
|
system_prompt=self.promot,
|
||||||
|
user_prompt=user_prompt
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.bind(tag=TAG).info(f"Detected intent: {intent}")
|
||||||
|
return intent.strip()
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
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, dialogue_history: List[Dict]) -> str:
|
||||||
|
"""
|
||||||
|
默认的意图识别实现,始终返回继续聊天
|
||||||
|
Args:
|
||||||
|
dialogue_history: 对话历史记录列表
|
||||||
|
Returns:
|
||||||
|
固定返回"继续聊天"
|
||||||
|
"""
|
||||||
|
logger.bind(tag=TAG).debug("Using NoIntentProvider, always returning continue chat")
|
||||||
|
return self.intent_options["continue_chat"]
|
||||||
@@ -1,8 +1,38 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
class LLMProviderBase(ABC):
|
class LLMProviderBase(ABC):
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def response(self, session_id, dialogue):
|
def response(self, session_id, dialogue):
|
||||||
"""LLM response generator"""
|
"""LLM response generator"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def response_no_stream(self, system_prompt, user_prompt):
|
||||||
|
try:
|
||||||
|
# 构造对话格式
|
||||||
|
dialogue = [
|
||||||
|
{"role": "system", "content": system_prompt},
|
||||||
|
{"role": "user", "content": user_prompt}
|
||||||
|
]
|
||||||
|
result = ""
|
||||||
|
for part in self.response("", dialogue):
|
||||||
|
result += part
|
||||||
|
return result
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error in Ollama response generation: {e}")
|
||||||
|
return "【LLM服务响应异常】"
|
||||||
|
|
||||||
|
def response_with_functions(self, session_id, dialogue, functions=None):
|
||||||
|
"""
|
||||||
|
Default implementation for function calling (streaming)
|
||||||
|
This should be overridden by providers that support function calls
|
||||||
|
|
||||||
|
Returns: generator that yields either text tokens or a special function call token
|
||||||
|
"""
|
||||||
|
# For providers that don't support functions, just return regular response
|
||||||
|
for token in self.response(session_id, dialogue):
|
||||||
|
yield {"type": "content", "content": token}
|
||||||
@@ -1,5 +1,6 @@
|
|||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
import requests, json
|
from openai import OpenAI
|
||||||
|
import json
|
||||||
from core.providers.llm.base import LLMProviderBase
|
from core.providers.llm.base import LLMProviderBase
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -8,39 +9,73 @@ logger = setup_logging()
|
|||||||
|
|
||||||
class LLMProvider(LLMProviderBase):
|
class LLMProvider(LLMProviderBase):
|
||||||
def __init__(self, config):
|
def __init__(self, config):
|
||||||
|
|
||||||
self.model_name = config.get("model_name")
|
self.model_name = config.get("model_name")
|
||||||
self.base_url = config.get("base_url", "http://localhost:11434")
|
self.base_url = config.get("base_url", "http://localhost:11434")
|
||||||
|
# Initialize OpenAI client with Ollama base URL
|
||||||
|
#如果没有v1,增加v1
|
||||||
|
if not self.base_url.endswith("/v1"):
|
||||||
|
self.base_url = f"{self.base_url}/v1"
|
||||||
|
|
||||||
|
self.client = OpenAI(
|
||||||
|
base_url=self.base_url,
|
||||||
|
api_key="ollama" # Ollama doesn't need an API key but OpenAI client requires one
|
||||||
|
)
|
||||||
|
|
||||||
def response(self, session_id, dialogue):
|
def response(self, session_id, dialogue):
|
||||||
try:
|
try:
|
||||||
# Convert dialogue format to Ollama format
|
responses = self.client.chat.completions.create(
|
||||||
prompt = ""
|
model=self.model_name,
|
||||||
for msg in dialogue:
|
messages=dialogue,
|
||||||
if msg["role"] == "system":
|
|
||||||
prompt += f"System: {msg['content']}\n"
|
|
||||||
elif msg["role"] == "user":
|
|
||||||
prompt += f"User: {msg['content']}\n"
|
|
||||||
elif msg["role"] == "assistant":
|
|
||||||
prompt += f"Assistant: {msg['content']}\n"
|
|
||||||
|
|
||||||
# Make request to Ollama API
|
|
||||||
response = requests.post(
|
|
||||||
f"{self.base_url}/api/generate",
|
|
||||||
json={
|
|
||||||
"model": self.model_name,
|
|
||||||
"prompt": prompt,
|
|
||||||
"stream": True
|
|
||||||
},
|
|
||||||
stream=True
|
stream=True
|
||||||
)
|
)
|
||||||
|
|
||||||
for line in response.iter_lines():
|
for chunk in responses:
|
||||||
if line:
|
try:
|
||||||
json_response = json.loads(line)
|
delta = chunk.choices[0].delta if getattr(chunk, 'choices', None) else None
|
||||||
if "response" in json_response:
|
content = delta.content if hasattr(delta, 'content') else ''
|
||||||
yield json_response["response"]
|
if content:
|
||||||
|
yield content
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error processing chunk: {e}")
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Error in Ollama response generation: {e}")
|
logger.bind(tag=TAG).error(f"Error in Ollama response generation: {e}")
|
||||||
yield "【Ollama服务响应异常】"
|
yield "【Ollama服务响应异常】"
|
||||||
|
|
||||||
|
def response_with_functions(self, session_id, dialogue, functions=None):
|
||||||
|
try:
|
||||||
|
stream = self.client.chat.completions.create(
|
||||||
|
model=self.model_name,
|
||||||
|
messages=dialogue,
|
||||||
|
stream=True,
|
||||||
|
tools=functions,
|
||||||
|
)
|
||||||
|
|
||||||
|
current_function_call = None
|
||||||
|
current_content = ""
|
||||||
|
|
||||||
|
for chunk in stream:
|
||||||
|
delta = chunk.choices[0].delta
|
||||||
|
|
||||||
|
if delta.content:
|
||||||
|
current_content += delta.content
|
||||||
|
yield {"type": "content", "content": delta.content}
|
||||||
|
|
||||||
|
if delta.tool_calls:
|
||||||
|
tool_call = delta.tool_calls[0]
|
||||||
|
# Handle the function call data using proper attribute access
|
||||||
|
if not current_function_call:
|
||||||
|
current_function_call = {
|
||||||
|
"function": {
|
||||||
|
"name": tool_call.function.name,
|
||||||
|
"arguments": tool_call.function.arguments
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if current_function_call:
|
||||||
|
logger.bind(tag=TAG).debug(f"ollama Function call detected: {current_function_call}")
|
||||||
|
yield {"type": "function_call", "function_call": current_function_call}
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error in Ollama function call: {e}")
|
||||||
|
yield {"type": "content", "content": f"【Ollama服务响应异常: {str(e)}】"}
|
||||||
@@ -43,3 +43,41 @@ class LLMProvider(LLMProviderBase):
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
|
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
|
||||||
|
|
||||||
|
def response_with_functions(self, session_id, dialogue, functions=None):
|
||||||
|
try:
|
||||||
|
stream = self.client.chat.completions.create(
|
||||||
|
model=self.model_name,
|
||||||
|
messages=dialogue,
|
||||||
|
stream=True,
|
||||||
|
tools=functions,
|
||||||
|
)
|
||||||
|
|
||||||
|
current_function_call = None
|
||||||
|
current_content = ""
|
||||||
|
|
||||||
|
for chunk in stream:
|
||||||
|
delta = chunk.choices[0].delta
|
||||||
|
|
||||||
|
if delta.content:
|
||||||
|
current_content += delta.content
|
||||||
|
yield {"type": "content", "content": delta.content}
|
||||||
|
|
||||||
|
if delta.tool_calls:
|
||||||
|
tool_call = delta.tool_calls[0]
|
||||||
|
# Handle the function call data using proper attribute access
|
||||||
|
if not current_function_call:
|
||||||
|
current_function_call = {
|
||||||
|
"function": {
|
||||||
|
"name": tool_call.function.name,
|
||||||
|
"arguments": tool_call.function.arguments
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if current_function_call:
|
||||||
|
logger.bind(tag=TAG).debug(f"openai Function call detected: {current_function_call}")
|
||||||
|
yield {"type": "function_call", "function_call": current_function_call}
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"Error in function call streaming: {e}")
|
||||||
|
yield {"type": "content", "content": f"【OpenAI服务响应异常: {e}】"}
|
||||||
@@ -8,6 +8,7 @@ class MemoryProviderBase(ABC):
|
|||||||
def __init__(self, config):
|
def __init__(self, config):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.role_id = None
|
self.role_id = None
|
||||||
|
self.llm = None
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def save_memory(self, msgs):
|
async def save_memory(self, msgs):
|
||||||
@@ -19,5 +20,6 @@ class MemoryProviderBase(ABC):
|
|||||||
"""Query memories for specific role based on similarity"""
|
"""Query memories for specific role based on similarity"""
|
||||||
return "please implement query method"
|
return "please implement query method"
|
||||||
|
|
||||||
def set_role_id(self, role_id: str):
|
def init_memory(self, role_id, llm):
|
||||||
self.role_id = role_id
|
self.role_id = role_id
|
||||||
|
self.llm = llm
|
||||||
|
|||||||
@@ -0,0 +1,156 @@
|
|||||||
|
from ..base import MemoryProviderBase, logger
|
||||||
|
import time
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import yaml
|
||||||
|
from core.utils.util import get_project_dir
|
||||||
|
|
||||||
|
short_term_memory_prompt = """
|
||||||
|
# 时空记忆编织者
|
||||||
|
|
||||||
|
## 核心使命
|
||||||
|
构建可生长的动态记忆网络,在有限空间内保留关键信息的同时,智能维护信息演变轨迹
|
||||||
|
根据对话记录,总结user的重要信息,以便在未来的对话中提供更个性化的服务
|
||||||
|
|
||||||
|
## 记忆法则
|
||||||
|
### 1. 三维度记忆评估(每次更新必执行)
|
||||||
|
| 维度 | 评估标准 | 权重分 |
|
||||||
|
|------------|---------------------------|--------|
|
||||||
|
| 时效性 | 信息新鲜度(按对话轮次) | 40% |
|
||||||
|
| 情感强度 | 含💖标记/重复提及次数 | 35% |
|
||||||
|
| 关联密度 | 与其他信息的连接数量 | 25% |
|
||||||
|
|
||||||
|
### 2. 动态更新机制
|
||||||
|
**名字变更处理示例:**
|
||||||
|
原始记忆:"曾用名": ["张三"], "现用名": "张三丰"
|
||||||
|
触发条件:当检测到「我叫X」「称呼我Y」等命名信号时
|
||||||
|
操作流程:
|
||||||
|
1. 将旧名移入"曾用名"列表
|
||||||
|
2. 记录命名时间轴:"2024-02-15 14:32:启用张三丰"
|
||||||
|
3. 在记忆立方追加:「从张三到张三丰的身份蜕变」
|
||||||
|
|
||||||
|
### 3. 空间优化策略
|
||||||
|
- **信息压缩术**:用符号体系提升密度
|
||||||
|
- ✅"张三丰[北/软工/🐱]"
|
||||||
|
- ❌"北京软件工程师,养猫"
|
||||||
|
- **淘汰预警**:当总字数≥900时触发
|
||||||
|
1. 删除权重分<60且3轮未提及的信息
|
||||||
|
2. 合并相似条目(保留时间戳最近的)
|
||||||
|
|
||||||
|
## 记忆结构
|
||||||
|
输出格式必须为可解析的json字符串,不需要解释、注释和说明,保存记忆时仅从对话提取信息,不要混入示例内容
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"时空档案": {
|
||||||
|
"身份图谱": {
|
||||||
|
"现用名": "",
|
||||||
|
"特征标记": []
|
||||||
|
},
|
||||||
|
"记忆立方": [
|
||||||
|
{
|
||||||
|
"事件": "入职新公司",
|
||||||
|
"时间戳": "2024-03-20",
|
||||||
|
"情感值": 0.9,
|
||||||
|
"关联项": ["下午茶"],
|
||||||
|
"保鲜期": 30
|
||||||
|
}
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"关系网络": {
|
||||||
|
"高频话题": {"职场": 12},
|
||||||
|
"暗线联系": [""]
|
||||||
|
},
|
||||||
|
"待响应": {
|
||||||
|
"紧急事项": ["需立即处理的任务"],
|
||||||
|
"潜在关怀": ["可主动提供的帮助"]
|
||||||
|
},
|
||||||
|
"高光语录": [
|
||||||
|
"最打动人心的瞬间,强烈的情感表达,user的原话"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
"""
|
||||||
|
|
||||||
|
def extract_json_data(json_code):
|
||||||
|
start = json_code.find("```json")
|
||||||
|
# 从start开始找到下一个```结束
|
||||||
|
end = json_code.find("```", start+1)
|
||||||
|
#print("start:", start, "end:", end)
|
||||||
|
if start == -1 or end == -1:
|
||||||
|
try:
|
||||||
|
jsonData = json.loads(json_code)
|
||||||
|
return json_code
|
||||||
|
except Exception as e:
|
||||||
|
print("Error:", e)
|
||||||
|
return ""
|
||||||
|
jsonData = json_code[start+7:end]
|
||||||
|
return jsonData
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
|
||||||
|
class MemoryProvider(MemoryProviderBase):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__(config)
|
||||||
|
self.short_momery = ""
|
||||||
|
self.memory_path = get_project_dir() + 'data/.memory.yaml'
|
||||||
|
self.load_memory()
|
||||||
|
|
||||||
|
def init_memory(self, role_id, llm):
|
||||||
|
super().init_memory(role_id, llm)
|
||||||
|
self.load_memory()
|
||||||
|
|
||||||
|
def load_memory(self):
|
||||||
|
all_memory = {}
|
||||||
|
if os.path.exists(self.memory_path):
|
||||||
|
with open(self.memory_path, 'r', encoding='utf-8') as f:
|
||||||
|
all_memory = yaml.safe_load(f) or {}
|
||||||
|
if self.role_id in all_memory:
|
||||||
|
self.short_momery = all_memory[self.role_id]
|
||||||
|
|
||||||
|
def save_memory_to_file(self):
|
||||||
|
all_memory = {}
|
||||||
|
if os.path.exists(self.memory_path):
|
||||||
|
with open(self.memory_path, 'r', encoding='utf-8') as f:
|
||||||
|
all_memory = yaml.safe_load(f) or {}
|
||||||
|
all_memory[self.role_id] = self.short_momery
|
||||||
|
with open(self.memory_path, 'w', encoding='utf-8') as f:
|
||||||
|
yaml.dump(all_memory, f, allow_unicode=True)
|
||||||
|
|
||||||
|
async def save_memory(self, msgs):
|
||||||
|
if self.llm is None:
|
||||||
|
logger.bind(tag=TAG).error("LLM is not set for memory provider")
|
||||||
|
return None
|
||||||
|
|
||||||
|
if len(msgs) < 2:
|
||||||
|
return None
|
||||||
|
|
||||||
|
msgStr = ""
|
||||||
|
for msg in msgs:
|
||||||
|
if msg.role == "user":
|
||||||
|
msgStr += f"User: {msg.content}\n"
|
||||||
|
elif msg.role== "assistant":
|
||||||
|
msgStr += f"Assistant: {msg.content}\n"
|
||||||
|
if len(self.short_momery) > 0:
|
||||||
|
msgStr+="历史记忆:\n"
|
||||||
|
msgStr+=self.short_momery
|
||||||
|
|
||||||
|
#当前时间
|
||||||
|
time_str = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime())
|
||||||
|
msgStr += f"当前时间:{time_str}"
|
||||||
|
|
||||||
|
result = self.llm.response_no_stream(short_term_memory_prompt, msgStr)
|
||||||
|
|
||||||
|
json_str = extract_json_data(result)
|
||||||
|
try:
|
||||||
|
json_data = json.loads(json_str) # 检查json格式是否正确
|
||||||
|
self.short_momery = json_str
|
||||||
|
except Exception as e:
|
||||||
|
print("Error:", e)
|
||||||
|
|
||||||
|
self.save_memory_to_file()
|
||||||
|
logger.bind(tag=TAG).info(f"Save memory successful - Role: {self.role_id}")
|
||||||
|
|
||||||
|
return self.short_momery
|
||||||
|
|
||||||
|
async def query_memory(self, query: str)-> str:
|
||||||
|
return self.short_momery
|
||||||
@@ -0,0 +1,18 @@
|
|||||||
|
'''
|
||||||
|
不使用记忆,可以选择此模块
|
||||||
|
'''
|
||||||
|
from ..base import MemoryProviderBase, logger
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
|
||||||
|
class MemoryProvider(MemoryProviderBase):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__(config)
|
||||||
|
|
||||||
|
async def save_memory(self, msgs):
|
||||||
|
logger.bind(tag=TAG).debug("nomem mode: No memory saving is performed.")
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def query_memory(self, query: str)-> str:
|
||||||
|
logger.bind(tag=TAG).debug("nomem mode: No memory query is performed.")
|
||||||
|
return ""
|
||||||
@@ -26,6 +26,9 @@ class Dialogue:
|
|||||||
return dialogue
|
return dialogue
|
||||||
|
|
||||||
def get_llm_dialogue_with_memory(self, memory_str: str = None) -> List[Dict[str, str]]:
|
def get_llm_dialogue_with_memory(self, memory_str: str = None) -> List[Dict[str, str]]:
|
||||||
|
if memory_str is None or len(memory_str) == 0:
|
||||||
|
return self.get_llm_dialogue()
|
||||||
|
|
||||||
# 构建带记忆的对话
|
# 构建带记忆的对话
|
||||||
dialogue = []
|
dialogue = []
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,17 @@
|
|||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from config.logger import setup_logging
|
||||||
|
import importlib
|
||||||
|
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
def create_instance(class_name, *args, **kwargs):
|
||||||
|
# 创建intent实例
|
||||||
|
if os.path.exists(os.path.join('core', 'providers', 'intent', class_name, f'{class_name}.py')):
|
||||||
|
lib_name = f'core.providers.intent.{class_name}.{class_name}'
|
||||||
|
if lib_name not in sys.modules:
|
||||||
|
sys.modules[lib_name] = importlib.import_module(f'{lib_name}')
|
||||||
|
return sys.modules[lib_name].IntentProvider(*args, **kwargs)
|
||||||
|
|
||||||
|
raise ValueError(f"不支持的intent类型: {class_name},请检查该配置的type是否设置正确")
|
||||||
@@ -4,7 +4,7 @@ from config.logger import setup_logging
|
|||||||
from core.connection import ConnectionHandler
|
from core.connection import ConnectionHandler
|
||||||
from core.handle.musicHandler import MusicHandler
|
from core.handle.musicHandler import MusicHandler
|
||||||
from core.utils.util import get_local_ip
|
from core.utils.util import get_local_ip
|
||||||
from core.utils import asr, vad, llm, tts, memory
|
from core.utils import asr, vad, llm, tts, memory, intent
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -13,11 +13,11 @@ class WebSocketServer:
|
|||||||
def __init__(self, config: dict):
|
def __init__(self, config: dict):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.logger = setup_logging()
|
self.logger = setup_logging()
|
||||||
self._vad, self._asr, self._llm, self._tts, self._music, self._memory = self._create_processing_instances()
|
self._vad, self._asr, self._llm, self._tts, self._music, self._memory, self.intent = self._create_processing_instances()
|
||||||
self.active_connections = set() # 添加全局连接记录
|
self.active_connections = set() # 添加全局连接记录
|
||||||
|
|
||||||
def _create_processing_instances(self):
|
def _create_processing_instances(self):
|
||||||
memory_cls_name = self.config["selected_module"].get("Memory", "mem0ai") # 默认使用mem0ai
|
memory_cls_name = self.config["selected_module"].get("Memory", "nomem") # 默认使用nomem
|
||||||
has_memory_cfg = self.config.get("Memory") and memory_cls_name in self.config["Memory"]
|
has_memory_cfg = self.config.get("Memory") and memory_cls_name in self.config["Memory"]
|
||||||
memory_cfg = self.config["Memory"][memory_cls_name] if has_memory_cfg else {}
|
memory_cfg = self.config["Memory"][memory_cls_name] if has_memory_cfg else {}
|
||||||
|
|
||||||
@@ -52,6 +52,13 @@ class WebSocketServer:
|
|||||||
),
|
),
|
||||||
MusicHandler(self.config),
|
MusicHandler(self.config),
|
||||||
memory.create_instance(memory_cls_name, memory_cfg),
|
memory.create_instance(memory_cls_name, memory_cfg),
|
||||||
|
intent.create_instance(
|
||||||
|
self.config["selected_module"]["Intent"]
|
||||||
|
if not 'type' in self.config["Intent"][self.config["selected_module"]["Intent"]]
|
||||||
|
else
|
||||||
|
self.config["Intent"][self.config["selected_module"]["Intent"]]["type"],
|
||||||
|
self.config["Intent"][self.config["selected_module"]["Intent"]]
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
async def start(self):
|
async def start(self):
|
||||||
@@ -71,7 +78,7 @@ class WebSocketServer:
|
|||||||
async def _handle_connection(self, websocket):
|
async def _handle_connection(self, websocket):
|
||||||
"""处理新连接,每次创建独立的ConnectionHandler"""
|
"""处理新连接,每次创建独立的ConnectionHandler"""
|
||||||
# 创建ConnectionHandler时传入当前server实例
|
# 创建ConnectionHandler时传入当前server实例
|
||||||
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._music, self._memory)
|
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._music, self._memory, self.intent)
|
||||||
self.active_connections.add(handler)
|
self.active_connections.add(handler)
|
||||||
try:
|
try:
|
||||||
await handler.handle_connection(websocket)
|
await handler.handle_connection(websocket)
|
||||||
|
|||||||
Reference in New Issue
Block a user