mergin main,速度提升一下

This commit is contained in:
lizhongxiang
2025-03-20 09:39:25 +08:00
96 changed files with 3438 additions and 1830 deletions
+104 -50
View File
@@ -5,17 +5,20 @@ import time
import queue
import asyncio
import traceback
from config.logger import setup_logging
import threading
import websockets
from typing import Dict, Any
import plugins_func.loadplugins
from config.logger import setup_logging
from core.utils.dialogue import Message, Dialogue
from core.handle.textHandle import handleTextMessage
from core.utils.util import get_string_no_punctuation_or_emoji, extract_json_from_string
from core.utils.util import get_string_no_punctuation_or_emoji, extract_json_from_string, get_ip_info
from concurrent.futures import ThreadPoolExecutor, TimeoutError
from core.handle.sendAudioHandle import sendAudioMessage, sendAudioMessageStream
from core.handle.receiveAudioHandle import handleAudioMessage
from core.handle.intentHandler import Action, get_functions, handle_llm_function_call
from core.handle.functionHandler import FunctionHandler
from plugins_func.register import Action
from config.private_config import PrivateConfig
from core.auth import AuthMiddleware, AuthenticationError
from core.utils.auth_code_gen import AuthCodeGenerator
@@ -28,7 +31,7 @@ class TTSException(RuntimeError):
class ConnectionHandler:
def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _music, _memory, _intent):
def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _memory, _intent):
self.config = config
self.logger = setup_logging()
self.auth = AuthMiddleware(config)
@@ -37,6 +40,8 @@ class ConnectionHandler:
self.websocket = None
self.headers = None
self.client_ip = None
self.client_ip_info = {}
self.session_id = None
self.prompt = None
self.welcome_msg = None
@@ -94,7 +99,6 @@ class ConnectionHandler:
self.private_config = None
self.auth_code_gen = AuthCodeGenerator.get_instance()
self.is_device_verified = False # 添加设备验证状态标志
self.music_handler = _music
self.close_after_chat = False # 是否在聊天结束后关闭连接
self.use_function_call_mode = False
if self.config["selected_module"]["Intent"] == 'function_call':
@@ -105,8 +109,8 @@ class ConnectionHandler:
# 获取并验证headers
self.headers = dict(ws.request.headers)
# 获取客户端ip地址
client_ip = ws.remote_address[0]
self.logger.bind(tag=TAG).info(f"{client_ip} conn - Headers: {self.headers}")
self.client_ip = ws.remote_address[0]
self.logger.bind(tag=TAG).info(f"{self.client_ip} conn - Headers: {self.headers}")
# 进行认证
await self.auth.authenticate(self.headers)
@@ -150,6 +154,7 @@ class ConnectionHandler:
self.welcome_msg["session_id"] = self.session_id
await self.websocket.send(json.dumps(self.welcome_msg))
# 异步初始化
await self.loop.run_in_executor(None, self._initialize_components)
# tts 消化线程
@@ -190,12 +195,21 @@ class ConnectionHandler:
self.prompt = self.config["prompt"]
if self.private_config:
self.prompt = self.private_config.private_config.get("prompt", self.prompt)
# 赋予LLM时间观念
if "{date_time}" in self.prompt:
date_time = time.strftime("%Y-%m-%d %H:%M", time.localtime())
self.prompt = self.prompt.replace("{date_time}", date_time)
self.client_ip_info = get_ip_info(self.client_ip)
self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}")
self.prompt = self.prompt + f"\n我在:{self.client_ip_info}"
self.dialogue.put(Message(role="system", content=self.prompt))
self.func_handler = FunctionHandler(self.config)
def change_system_prompt(self, prompt):
self.prompt = prompt
# 找到原来的role==system,替换原来的系统提示
for m in self.dialogue.dialogue:
if m.role == "system":
m.content = prompt
async def _check_and_broadcast_auth_code(self):
"""检查设备绑定状态并广播认证码"""
if not self.private_config.get_owner():
@@ -312,7 +326,7 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False))
return True
def chat_with_function_calling(self, query):
def chat_with_function_calling(self, query, tool_call=False):
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():
@@ -321,10 +335,11 @@ class ConnectionHandler:
future.result()
return True
self.dialogue.put(Message(role="user", content=query))
if not tool_call:
self.dialogue.put(Message(role="user", content=query))
# Define intent functions
functions = get_functions()
functions = self.func_handler.get_functions()
response_message = []
processed_chars = 0 # 跟踪已处理的字符位置
@@ -360,7 +375,7 @@ class ConnectionHandler:
for response in llm_responses:
content, tools_call = response
if content is not None and len(content) > 0:
if len(response_message) <= 0 and content == "```":
if len(response_message) <= 0 and (content == "```" or "<tool_call>" in content):
tool_call_flag = True
if tools_call is not None:
@@ -417,6 +432,38 @@ class ConnectionHandler:
self.tts_queue.put(future)
processed_chars += len(segment_text_raw) # 更新已处理字符位置
# 处理function call
if tool_call_flag:
bHasError = False
if function_id is None:
a = extract_json_from_string(content_arguments)
if a is not None:
try:
content_arguments_json = json.loads(a)
function_name = content_arguments_json["name"]
function_arguments = json.dumps(content_arguments_json["arguments"], ensure_ascii=False)
function_id = str(uuid.uuid4().hex)
except Exception as e:
bHasError = True
response_message.append(a)
else:
bHasError = True
response_message.append(content_arguments)
if bHasError:
self.logger.bind(tag=TAG).error(f"function call error: {content_arguments}")
else:
function_arguments = json.loads(function_arguments)
if not bHasError:
self.logger.bind(tag=TAG).info(
f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}")
function_call_data = {
"name": function_name,
"id": function_id,
"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)
# 处理最后剩余的文本
full_text = "".join(response_message)
remaining_text = full_text[processed_chars:]
@@ -424,6 +471,7 @@ class ConnectionHandler:
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)
if self.tts_stream:
stream_queue = queue.Queue()
self.executor.submit(self.speak_and_play_stream, segment_text, stream_queue, text_index)
@@ -433,36 +481,13 @@ class ConnectionHandler:
"text_index": text_index
})
else:
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.tts_queue.put(future)
# 存储对话内容
if len(response_message) > 0:
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
# 处理function call
if tool_call_flag:
if function_id is None:
a = extract_json_from_string(content_arguments)
if a is not None:
content_arguments_json = json.loads(a)
function_name = content_arguments_json["function_name"]
function_arguments = json.dumps(content_arguments_json["args"], ensure_ascii=False)
function_id = str(uuid.uuid4().hex)
else:
return []
function_arguments = json.loads(function_arguments)
self.logger.bind(tag=TAG).info(
f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}")
function_call_data = {
"name": function_name,
"id": function_id,
"arguments": function_arguments
}
result = handle_llm_function_call(self, function_call_data)
self._handle_function_result(result, function_call_data, text_index + 1)
self.llm_finish_task = True
self.logger.bind(tag=TAG).debug(json.dumps(self.dialogue.get_llm_dialogue(), indent=4, ensure_ascii=False))
@@ -484,10 +509,34 @@ class ConnectionHandler:
future = self.executor.submit(self.speak_and_play, text, text_index)
self.tts_queue.put(future)
self.dialogue.put(Message(role="assistant", content=text))
if result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
text = result.response
if result.action == Action.NOTFOUND:
text = result.response
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
text = result.result
if text is not None and len(text) > 0:
function_id = function_call_data["id"]
function_name = function_call_data["name"]
function_arguments = function_call_data["arguments"]
self.dialogue.put(Message(role='assistant',
tool_calls=[{"id": function_id,
"function": {"arguments": function_arguments,
"name": function_name},
"type": 'function',
"index": 0}]))
self.dialogue.put(Message(role="tool", tool_call_id=function_id, content=text))
self.chat_with_function_calling(text, tool_call=True)
elif result.action == Action.NOTFOUND:
text = result.result
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.dialogue.put(Message(role="assistant", content=text))
else:
text = result.result
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.dialogue.put(Message(role="assistant", content=text))
def _tts_priority_thread(self):
if self.tts_stream:
@@ -506,7 +555,8 @@ class ConnectionHandler:
opus_datas, text_index, tts_file = [], 0, None
try:
self.logger.bind(tag=TAG).debug("正在处理TTS任务...")
tts_file, text, text_index = future.result(timeout=10)
tts_timeout = self.config.get("tts_timeout", 10)
tts_file, text, text_index = future.result(timeout=tts_timeout)
if text is None or len(text) <= 0:
self.logger.bind(tag=TAG).error(f"TTS出错:{text_index}: tts text is empty")
elif tts_file is None:
@@ -514,7 +564,7 @@ class ConnectionHandler:
else:
self.logger.bind(tag=TAG).debug(f"TTS生成:文件路径: {tts_file}")
if os.path.exists(tts_file):
opus_datas, duration = self.tts.wav_to_opus_data(tts_file)
opus_datas, duration = self.tts.audio_to_opus_data(tts_file)
else:
self.logger.bind(tag=TAG).error(f"TTS出错:文件不存在{tts_file}")
except TimeoutError:
@@ -596,17 +646,21 @@ class ConnectionHandler:
return tts_file, text, text_index
def speak_and_play_stream(self, text, queue: queue.Queue, text_index=0):
if text is None or len(text) <= 0:
self.logger.bind(tag=TAG).info(f"无需tts转换,query为空,{text}")
return None, text
self.tts.to_tts_stream(text, queue, text_index)
try:
if text is None or len(text) <= 0:
self.logger.bind(tag=TAG).info(f"无需tts转换,query为空,{text}")
return None, text
self.tts.to_tts_stream(text, queue, text_index)
except Exception as e:
self.logger.bind(tag=TAG).error(e)
traceback.print_exc()
raise e
def clearSpeakStatus(self):
self.logger.bind(tag=TAG).debug(f"清除服务端讲话状态")
self.asr_server_receive = True
self.tts_last_text_index = -1
self.tts_first_text_index = -1
self.tts_duration = 0
def recode_first_last_text(self, text, text_index=0):
if self.tts_first_text_index == -1:
@@ -0,0 +1,84 @@
import asyncio
from enum import Enum
from config.logger import setup_logging
import json
from plugins_func.register import FunctionRegistry, ActionResponse, Action, ToolType
TAG = __name__
logger = setup_logging()
class FunctionHandler:
def __init__(self, config):
self.config = config
self.function_registry = FunctionRegistry()
self.register_nessary_functions()
self.register_config_functions()
self.functions_desc = self.function_registry.get_all_function_desc()
func_names = self.current_support_functions()
self.modify_plugin_loader_des(func_names)
def modify_plugin_loader_des(self, func_names):
if "plugin_loader" not in func_names:
return
# 可编辑的列表中去掉plugin_loader
surport_plugins = [func for func in func_names if func != "plugin_loader"]
func_names = ",".join(surport_plugins)
for function_desc in self.functions_desc:
if function_desc["function"]["name"] == "plugin_loader":
function_desc["function"]["description"] = function_desc["function"]["description"].replace("[plugins]", func_names)
break
def upload_functions_desc(self):
self.functions_desc = self.function_registry.get_all_function_desc()
def current_support_functions(self):
func_names = []
for func in self.functions_desc:
func_names.append(func["function"]["name"])
# 打印当前支持的函数列表
logger.bind(tag=TAG).info(f"当前支持的函数列表: {func_names}")
return func_names
def get_functions(self):
"""获取功能调用配置"""
return self.functions_desc
def register_nessary_functions(self):
"""注册必要的函数"""
self.function_registry.register_function("handle_exit_intent")
self.function_registry.register_function("play_music")
self.function_registry.register_function("plugin_loader")
self.function_registry.register_function("get_time")
self.function_registry.register_function("raise_and_lower_the_volume")
def register_config_functions(self):
"""注册配置中的函数,可以不同客户端使用不同的配置"""
for func in self.config["Intent"]["function_call"].get("functions", []):
self.function_registry.register_function(func)
def get_function(self, name):
return self.function_registry.get_function(name)
def handle_llm_function_call(self, conn, function_call_data):
try:
function_name = function_call_data["name"]
funcItem = self.get_function(function_name)
if not funcItem:
return ActionResponse(action=Action.NOTFOUND, result="没有找到对应的函数", response="")
func = funcItem.func
arguments = function_call_data["arguments"]
arguments = json.loads(arguments) if arguments else {}
logger.bind(tag=TAG).info(f"调用函数: {function_name}, 参数: {arguments}")
if funcItem.type == ToolType.SYSTEM_CTL or funcItem.type == ToolType.IOT_CTL:
return func(conn, **arguments)
elif funcItem.type == ToolType.WAIT:
return func(**arguments)
elif funcItem.type == ToolType.CHANGE_SYS_PROMPT:
return func(conn, **arguments)
else:
return ActionResponse(action=Action.NOTFOUND, result="没有找到对应的函数", response="")
except Exception as e:
logger.bind(tag=TAG).error(f"处理function call错误: {e}")
return None
@@ -1,105 +1,24 @@
from config.logger import setup_logging
import json
import uuid
from core.handle.sendAudioHandle import send_stt_message
from core.utils.dialogue import Message
from core.utils.util import remove_punctuation_and_length
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
# 使用LLM进行意图分析
intent = await analyze_intent_with_llm(conn, text)
if not intent:
return False
# 处理各种意图
return await process_intent_result(conn, intent, text)
@@ -126,7 +45,6 @@ async def analyze_intent_with_llm(conn, text):
dialogue = conn.dialogue
try:
intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text)
# 尝试解析JSON结果
try:
intent_data = json.loads(intent_result)
@@ -147,9 +65,6 @@ 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)
@@ -158,10 +73,37 @@ async def process_intent_result(conn, intent, original_text):
# 处理播放音乐意图
if "播放音乐" in intent:
logger.bind(tag=TAG).info(f"识别到音乐播放意图: {intent}")
await conn.music_handler.handle_music_command(conn, intent)
# 调用play_music函数来播放音乐
song_name = extract_text_in_brackets(intent)
function_id = str(uuid.uuid4().hex)
function_name = "play_music"
function_arguments = '{ "song_name": "' + song_name + '" }'
function_call_data = {
"name": function_name,
"id": function_id,
"arguments": function_arguments
}
conn.func_handler.handle_llm_function_call(conn, function_call_data)
return True
# 其他意图处理可以在这里扩展
# 默认返回False,表示继续常规聊天流程
return False
def extract_text_in_brackets(s):
"""
从字符串中提取中括号内的文字
:param s: 输入字符串
:return: 中括号内的文字,如果不存在则返回空字符串
"""
left_bracket_index = s.find('[')
right_bracket_index = s.find(']')
if left_bracket_index != -1 and right_bracket_index != -1 and left_bracket_index < right_bracket_index:
return s[left_bracket_index + 1:right_bracket_index]
else:
return ""
+242 -97
View File
@@ -1,24 +1,120 @@
import json
import asyncio
from config.logger import setup_logging
from plugins_func.register import device_type_registry, register_function, ActionResponse, Action, ToolType
TAG = __name__
logger = setup_logging()
def wrap_async_function(async_func):
"""包装异步函数为同步函数"""
def wrapper(*args, **kwargs):
try:
# 获取连接对象(第一个参数)
conn = args[0]
if not hasattr(conn, 'loop'):
logger.bind(tag=TAG).error("Connection对象没有loop属性")
return ActionResponse(Action.ERROR, "Connection对象没有loop属性",
"执行操作时出错: Connection对象没有loop属性")
# 使用conn对象中的事件循环
loop = conn.loop
# 在conn的事件循环中运行异步函数
future = asyncio.run_coroutine_threadsafe(async_func(*args, **kwargs), loop)
# 等待结果返回
return future.result()
except Exception as e:
logger.bind(tag=TAG).error(f"运行异步函数时出错: {e}")
return ActionResponse(Action.ERROR, str(e), f"执行操作时出错: {e}")
return wrapper
def create_iot_function(device_name, method_name, method_info):
"""
根据IOT设备描述生成通用的控制函数
"""
async def iot_control_function(conn, response_success=None, response_failure=None, **params):
try:
# 打印响应参数
logger.bind(tag=TAG).info(
f"控制函数接收到的响应参数: success='{response_success}', failure='{response_failure}'")
# 发送控制命令
await send_iot_conn(conn, device_name, method_name, params)
# 等待一小段时间让状态更新
await asyncio.sleep(0.1)
# 生成结果信息
result = f"{device_name}{method_name}操作执行成功"
# 处理响应中可能的占位符
response = response_success
# 替换{value}占位符
for param_name, param_value in params.items():
# 先尝试直接替换参数值
if "{" + param_name + "}" in response:
response = response.replace("{" + param_name + "}", str(param_value))
# 如果有{value}占位符,用相关参数替换
if "{value}" in response:
response = response.replace("{value}", str(param_value))
break
return ActionResponse(Action.RESPONSE, result, response)
except Exception as e:
logger.bind(tag=TAG).error(f"执行{device_name}{method_name}操作失败: {e}")
# 操作失败时使用大模型提供的失败响应
response = response_failure
return ActionResponse(Action.ERROR, str(e), response)
return wrap_async_function(iot_control_function)
def create_iot_query_function(device_name, prop_name, prop_info):
"""
根据IOT设备属性创建查询函数
"""
async def iot_query_function(conn, response_success=None, response_failure=None):
try:
# 打印响应参数
logger.bind(tag=TAG).info(
f"查询函数接收到的响应参数: success='{response_success}', failure='{response_failure}'")
value = await get_iot_status(conn, device_name, prop_name)
# 查询成功,生成结果
if value is not None:
# 使用大模型提供的成功响应,并替换其中的占位符
response = response_success.replace("{value}", str(value))
return ActionResponse(Action.RESPONSE, str(value), response)
else:
# 查询失败,使用大模型提供的失败响应
response = response_failure
return ActionResponse(Action.ERROR, f"属性{prop_name}不存在", response)
except Exception as e:
logger.bind(tag=TAG).error(f"查询{device_name}{prop_name}时出错: {e}")
# 查询出错时使用大模型提供的失败响应
response = response_failure
return ActionResponse(Action.ERROR, str(e), response)
return wrap_async_function(iot_query_function)
class IotDescriptor:
"""
A class to represent an IoT descriptor.
Attributes:
----------
name : str
The name of the IoT descriptor.
description : str
A brief description of the IoT descriptor.
properties : dict
A dictionary containing properties of the IoT descriptor.
methods : dict
A dictionary containing methods of the IoT descriptor.
-------
"""
def __init__(self, name, description, properties, methods):
@@ -29,17 +125,7 @@ class IotDescriptor:
# 根据描述创建属性
for key, value in properties.items():
# "volume":{"description":"当前音量 值","type":"number"}
"""
等价于
{
'name': 名字,
'description': 描述,
'value': 0
}
"""
# setattr(self, key, {}) # 创建一个空字典, 名字是属性名
property_item = globals()[key] = {} # 创建一个空字典, 名字是属性名
property_item = globals()[key] = {}
property_item['name'] = key
property_item["description"] = value["description"]
if value["type"] == "number":
@@ -52,23 +138,10 @@ class IotDescriptor:
# 根据描述创建方法
for key, value in methods.items():
# "SetVolume": {"description":"设置音量","parameters":{"volume":{"description":"0到100之间的整数","type":"number"}}}
"""
等价于
SetVolume = {
`description`: 描述,
`volume`: {
`description`: 描述,
`value`: 0
}
}
"""
# setattr(self, key, {}) # 创建一个空字典, 名字是方法名
method = globals()[key] = {} # 创建一个空字典, 名字是方法名
method = globals()[key] = {}
method["description"] = value["description"]
method['name'] = key
for k, v in value["parameters"].items():
# 不同的参数解析
method[k] = {}
method[k]["description"] = v["description"]
if v["type"] == "number":
@@ -77,58 +150,136 @@ class IotDescriptor:
method[k]["value"] = False
else:
method[k]["value"] = ""
self.methods.append(method)
async def handleIotDescriptors(conn, descriptors):
"""
处理物联网描述
示例: [{
"name":"Speaker",
"description":"当前 AI 机器人的扬声器",
"properties":{
"volume":{"description":"当前音量 值","type":"number"} 可以有boolean, number, string三种类型
},
"methods":{
"SetVolume":{
"description":"设置音量","parameters":{"volume":{"description":"0到100之间的整数","type":"number"}}
def register_device_type(descriptor):
"""注册设备类型及其功能"""
device_name = descriptor["name"]
type_id = device_type_registry.generate_device_type_id(descriptor)
# 如果该类型已注册,直接返回类型ID
if type_id in device_type_registry.type_functions:
return type_id
functions = {}
# 为每个属性创建查询函数
for prop_name, prop_info in descriptor["properties"].items():
func_name = f"get_{device_name.lower()}_{prop_name.lower()}"
func_desc = {
"type": "function",
"function": {
"name": func_name,
"description": f"查询{descriptor['description']}{prop_info['description']}",
"parameters": {
"type": "object",
"properties": {
"response_success": {
"type": "string",
"description": f"查询成功时的友好回复,必须使用{{value}}作为占位符表示查询到的值"
},
"response_failure": {
"type": "string",
"description": f"查询失败时的友好回复,例如:'无法获取{device_name}{prop_info['description']}'"
}
},
"required": ["response_success", "response_failure"]
}
}
}
}]
descriptors: 描述列表
"""
query_func = create_iot_query_function(device_name, prop_name, prop_info)
decorated_func = register_function(func_name, func_desc, ToolType.IOT_CTL)(query_func)
functions[func_name] = decorated_func
# 为每个方法创建控制函数
for method_name, method_info in descriptor["methods"].items():
func_name = f"{device_name.lower()}_{method_name.lower()}"
# 创建参数字典,添加原有参数
parameters = {
param_name: {
"type": param_info["type"],
"description": param_info["description"]
}
for param_name, param_info in method_info["parameters"].items()
}
# 添加响应参数
parameters.update({
"response_success": {
"type": "string",
"description": "操作成功时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称"
},
"response_failure": {
"type": "string",
"description": "操作失败时的友好回复,关于该设备的操作结果,设备名称尽量使用description中的名称"
}
})
# 构建必须参数列表(原有参数 + 响应参数)
required_params = list(method_info["parameters"].keys())
required_params.extend(["response_success", "response_failure"])
func_desc = {
"type": "function",
"function": {
"name": func_name,
"description": f"{descriptor['description']} - {method_info['description']}",
"parameters": {
"type": "object",
"properties": parameters,
"required": required_params
}
}
}
control_func = create_iot_function(device_name, method_name, method_info)
decorated_func = register_function(func_name, func_desc, ToolType.IOT_CTL)(control_func)
functions[func_name] = decorated_func
device_type_registry.register_device_type(type_id, functions)
return type_id
# 用于接受前端设备推送的搜索iot描述
async def handleIotDescriptors(conn, descriptors):
"""处理物联网描述"""
functions_changed = False
for descriptor in descriptors:
# 创建IOT设备描述符
iot_descriptor = IotDescriptor(descriptor["name"], descriptor["description"], descriptor["properties"],
descriptor["methods"])
conn.iot_descriptors[descriptor["name"]] = iot_descriptor
# 暂时从配置文件中设置音量,后期通过意图识别控制音量
default_iot_volume = 100
if "iot" in conn.config:
default_iot_volume = conn.config["iot"]["Speaker"]["volume"]
logger.bind(tag=TAG).info(f"服务端设置音量为{default_iot_volume}")
await send_iot_conn(conn, "Speaker", "SetVolume", {"volume": default_iot_volume})
if conn.use_function_call_mode:
# 注册或获取设备类型
type_id = register_device_type(descriptor)
device_functions = device_type_registry.get_device_functions(type_id)
# 在连接级注册设备函数
if hasattr(conn, 'func_handler'):
for func_name in device_functions:
conn.func_handler.function_registry.register_function(func_name)
logger.bind(tag=TAG).info(f"注册IOT函数到function handler: {func_name}")
functions_changed = True
# 如果注册了新函数,更新function描述列表
if functions_changed and hasattr(conn, 'func_handler'):
conn.func_handler.upload_functions_desc()
func_names = conn.func_handler.current_support_functions()
logger.bind(tag=TAG).info(f"设备类型: {type_id}")
logger.bind(tag=TAG).info(f"更新function描述列表完成,当前支持的函数: {func_names}")
async def handleIotStatus(conn, states):
"""
处理物联网状态
示例: [{
"name":"Speaker",
"state":{
"volume":100
}
}]
states: 状态列表
"""
"""处理物联网状态"""
for state in states:
for key, value in conn.iot_descriptors.items():
if key == state["name"]:
for property_item in value.properties:
# properties为字典列表, 记录各种属性
for k, v in state["state"].items():
# state为字典, 记录各种属性的值, 是需要记录的信息
if property_item["name"] == k:
# 检查一下属性是不是相同的
if type(v) != type(property_item["value"]):
logger.bind(tag=TAG).error(f"属性{property_item['name']}的值类型不匹配")
break
@@ -138,41 +289,35 @@ async def handleIotStatus(conn, states):
break
break
async def get_iot_status(conn, name, property_name):
"""
获取物联网状态
name: 设备名称 "Speaker"
property_name: 属性名称 "volume"
返回值: 属性值, 实际的属性有int, bool和str三种类型
"""
"""获取物联网状态"""
for key, value in conn.iot_descriptors.items():
if key == name:
for property_item in value.properties:
if property_item["name"] == property_name:
return property_item["value"]
logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
return None
async def send_iot_conn(conn, name, method_name, parameters):
"""
发送物联网指令
name: 设备名称 "Speaker"
method: 方法 "SetVolume"
parameters: 参数, 是一个字典 {"volume": 100}
发送示例:
{
"type": "iot",
"commands": [
{
"name" : "Speaker",
"method": "SetVolume",
"parameters": {
"volume": 100
}
}
]
}
"""
async def set_iot_status(conn, name, property_name, value):
"""设置物联网状态"""
for key, iot_descriptor in conn.iot_descriptors.items():
if key == name:
for property_item in iot_descriptor.properties:
if property_item["name"] == property_name:
if type(value) != type(property_item["value"]):
logger.bind(tag=TAG).error(f"属性{property_item['name']}的值类型不匹配")
return
property_item["value"] = value
logger.bind(tag=TAG).info(f"物联网状态更新: {name} , {property_name} = {value}")
return
logger.bind(tag=TAG).warning(f"未找到设备 {name} 的属性 {property_name}")
async def send_iot_conn(conn, name, method_name, parameters):
"""发送物联网指令"""
for key, value in conn.iot_descriptors.items():
if key == name:
# 找到了设备
@@ -1,139 +0,0 @@
from config.logger import setup_logging
import os
import random
import difflib
import re
import traceback
from pathlib import Path
import time
from core.handle.sendAudioHandle import send_stt_message
from core.utils import p3
TAG = __name__
logger = setup_logging()
def _extract_song_name(text):
"""从用户输入中提取歌名"""
for keyword in ["播放音乐"]:
if keyword in text:
parts = text.split(keyword)
if len(parts) > 1:
return parts[1].strip()
return None
def _find_best_match(potential_song, music_files):
"""查找最匹配的歌曲"""
best_match = None
highest_ratio = 0
for music_file in music_files:
song_name = os.path.splitext(music_file)[0]
ratio = difflib.SequenceMatcher(None, potential_song, song_name).ratio()
if ratio > highest_ratio and ratio > 0.4:
highest_ratio = ratio
best_match = music_file
return best_match
class MusicManager:
def __init__(self, music_dir, music_ext):
self.music_dir = Path(music_dir)
self.music_ext = music_ext
def get_music_files(self):
music_files = []
for file in self.music_dir.rglob("*"):
# 判断是否是文件
if file.is_file():
# 获取文件扩展名
ext = file.suffix.lower()
# 判断扩展名是否在列表中
if ext in self.music_ext:
# music_files.append(str(file.resolve())) # 添加绝对路径
# 添加相对路径
music_files.append(str(file.relative_to(self.music_dir)))
return music_files
class MusicHandler:
def __init__(self, config):
self.config = config
if "music" in self.config:
self.music_config = self.config["music"]
self.music_dir = os.path.abspath(
self.music_config.get("music_dir", "./music") # 默认路径修改
)
self.music_ext = self.music_config.get("music_ext", (".mp3", ".wav", ".p3"))
self.refresh_time = self.music_config.get("refresh_time", 60)
else:
self.music_dir = os.path.abspath("./music")
self.music_ext = (".mp3", ".wav", ".p3")
self.refresh_time = 60
# 获取音乐文件列表
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}")
async def handle_music_command(self, conn, text):
"""处理音乐播放指令"""
clean_text = re.sub(r'[^\w\s]', '', text).strip()
logger.bind(tag=TAG).debug(f"检查是否是音乐命令: {clean_text}")
# 尝试匹配具体歌名
if os.path.exists(self.music_dir):
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}")
potential_song = _extract_song_name(clean_text)
if potential_song:
best_match = _find_best_match(potential_song, self.music_files)
if best_match:
logger.bind(tag=TAG).info(f"找到最匹配的歌曲: {best_match}")
await self.play_local_music(conn, specific_file=best_match)
return True
# 检查是否是通用播放音乐命令
await self.play_local_music(conn)
return True
async def play_local_music(self, conn, specific_file=None):
"""播放本地音乐文件"""
try:
if not os.path.exists(self.music_dir):
logger.bind(tag=TAG).error(f"音乐目录不存在: {self.music_dir}")
return
# 确保路径正确性
if specific_file:
selected_music = specific_file
music_path = os.path.join(self.music_dir, specific_file)
else:
if not self.music_files:
logger.bind(tag=TAG).error("未找到MP3音乐文件")
return
selected_music = random.choice(self.music_files)
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}")
return
text = f"正在播放{selected_music}"
await send_stt_message(conn, text)
conn.tts_first_text_index = 0
conn.tts_last_text_index = 0
conn.llm_finish_task = True
if music_path.endswith(".p3"):
opus_packets, duration = p3.decode_opus_from_file(music_path)
else:
opus_packets, duration = conn.tts.wav_to_opus_data(music_path)
conn.audio_play_queue.put((opus_packets, selected_music, 0))
except Exception as e:
logger.bind(tag=TAG).error(f"播放音乐失败: {str(e)}")
logger.bind(tag=TAG).error(f"详细错误: {traceback.format_exc()}")
@@ -20,7 +20,8 @@ async def handleAudioMessage(conn, audio):
# 如果本次没有声音,本段也没声音,就把声音丢弃了
if have_voice == False and conn.client_have_voice == False:
await no_voice_close_connect(conn)
conn.asr_audio.clear()
conn.asr_audio.append(audio)
conn.asr_audio = conn.asr_audio[-5:] # 保留最新的5帧音频内容,解决ASR句首丢字问题
return
conn.client_no_voice_last_time = 0.0
conn.asr_audio.append(audio)
@@ -29,7 +30,7 @@ async def handleAudioMessage(conn, audio):
conn.client_abort = False
conn.asr_server_receive = False
# 音频太短了,无法识别
if len(conn.asr_audio) < 3:
if len(conn.asr_audio) < 10:
conn.asr_server_receive = True
else:
text, file_path = await conn.asr.speech_to_text(conn.asr_audio, conn.session_id)
@@ -51,15 +51,6 @@ async def sendAudioMessageStream(conn, audios_queue, text, text_index=0, llm_fin
for opus_packet in audio_opus_datas:
if conn.client_abort:
return
# 计算当前包的预期发送时间
# 计算当前包的预期发送时间
expected_time = start_time_chunk + (play_position / 1000)
current_time = time.perf_counter()
# 等待直到预期时间
delay = expected_time - current_time
if delay > 0:
await asyncio.sleep(delay)
logger.bind(tag=TAG).info(f'发送数据长度:{len(opus_packet)}')
await conn.websocket.send(opus_packet)
play_position += frame_duration # 更新播放位置
@@ -70,15 +61,15 @@ async def sendAudioMessageStream(conn, audios_queue, text, text_index=0, llm_fin
await send_tts_message(conn, "sentence_end", text)
print(f'{text_index}-{conn.tts_last_text_index}')
expected_time = start_time_chunk + (play_position / 1000)
current_time = time.perf_counter()
# 等待直到预期时间
delay = expected_time - current_time
if delay > 0:
await asyncio.sleep(delay)
# 发送结束消息(如果是最后一个文本)
logger.bind(tag=TAG).info(f"{conn.llm_finish_task},{text_index},{conn.tts_last_text_index}")
if conn.llm_finish_task and text_index == conn.tts_last_text_index:
expected_time = start_time_chunk + (play_position / 1000)
current_time = time.perf_counter()
# 等待直到预期时间
delay = expected_time - current_time
if delay > 0:
await asyncio.sleep(delay)
await send_tts_message(conn, 'stop', None)
if conn.close_after_chat or "拜拜" in text or "再见" in text:
await conn.close()
@@ -2,7 +2,7 @@ from config.logger import setup_logging
import json
from core.handle.abortHandle import handleAbortMessage
from core.handle.helloHandle import handleHelloMessage
from core.handle.receiveAudioHandle import startToChat
from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
from core.handle.iotHandle import handleIotDescriptors, handleIotStatus
TAG = __name__
@@ -31,6 +31,8 @@ async def handleTextMessage(conn, message):
elif msg_json["state"] == "stop":
conn.client_have_voice = True
conn.client_voice_stop = True
if len(conn.asr_audio) > 0:
await handleAudioMessage(conn, b'')
elif msg_json["state"] == "detect":
conn.asr_server_receive = False
conn.client_have_voice = False
@@ -41,6 +43,6 @@ async def handleTextMessage(conn, message):
if "descriptors" in msg_json:
await handleIotDescriptors(conn, msg_json["descriptors"])
if "states" in msg_json:
await handleIotStatus(conn, msg_json["states"])
await handleIotStatus(conn, msg_json["states"])
except json.JSONDecodeError:
await conn.websocket.send(message)
@@ -1,12 +1,13 @@
from typing import List, Dict
from ..base import IntentProviderBase
from plugins_func.functions.play_music import initialize_music_handler
from config.logger import setup_logging
import re
import re
TAG = __name__
logger = setup_logging()
class IntentProvider(IntentProviderBase):
def __init__(self, config):
super().__init__(config)
@@ -73,8 +74,8 @@ class IntentProvider(IntentProviderBase):
"你现在可以使用的音乐的名称如下(使用<start>和<end>标志):\n"
)
return prompt
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:
raise ValueError("LLM provider not set")
@@ -89,7 +90,9 @@ class IntentProvider(IntentProviderBase):
msgStr += f"User: {text}\n"
user_prompt = f"当前的对话如下:\n{msgStr}"
prompt_music = f"{self.promot}\n<start>{conn.music_handler.music_files}\n<end>"
music_config = initialize_music_handler(conn)
music_file_names = music_config["music_file_names"]
prompt_music = f"{self.promot}\n<start>{music_file_names}\n<end>"
logger.bind(tag=TAG).debug(f"User prompt: {prompt_music}")
# 使用LLM进行意图识别
intent = self.llm.response_no_stream(
@@ -100,10 +103,9 @@ class IntentProvider(IntentProviderBase):
# 使用正则表达式提取 {} 中的内容
match = re.search(r'\{.*?\}', intent)
if match:
result = match.group(0) # 获取匹配到的内容(包含 {}
print(result) # 输出:{intent: '播放音乐 [中秋月]'}
result = match.group(0)
intent = result
else:
intent = "{intent: '继续聊天'}"
logger.bind(tag=TAG).info(f"Detected intent: {intent}")
return intent.strip()
return intent.strip()
@@ -10,11 +10,13 @@ class LLMProvider(LLMProviderBase):
def __init__(self, config):
self.api_key = config["api_key"]
self.base_url = config.get("base_url", "https://api.dify.ai/v1").rstrip('/')
self.session_conversation_map = {} # 存储session_id和conversation_id的映射
def response(self, session_id, dialogue):
try:
# 取最后一条用户消息
last_msg = next(m for m in reversed(dialogue) if m["role"] == "user")
conversation_id = self.session_conversation_map.get(session_id)
# 发起流式请求
with requests.post(
@@ -24,13 +26,18 @@ class LLMProvider(LLMProviderBase):
"query": last_msg["content"],
"response_mode": "streaming",
"user": session_id,
"inputs": {}
"inputs": {},
"conversation_id": conversation_id
},
stream=True
) as r:
for line in r.iter_lines():
if line.startswith(b'data: '):
event = json.loads(line[6:])
# 如果没有找到conversation_id,则获取此次conversation_id
if not conversation_id:
conversation_id = event.get('conversation_id')
self.session_conversation_map[session_id] = conversation_id # 更新映射
if event.get('answer'):
yield event['answer']
@@ -12,10 +12,10 @@ class LLMProvider(LLMProviderBase):
self.model_name = config.get("model_name")
self.base_url = config.get("base_url", "http://localhost:11434")
# Initialize OpenAI client with Ollama base URL
#如果没有v1,增加v1
# 如果没有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
@@ -28,13 +28,20 @@ class LLMProvider(LLMProviderBase):
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:
yield 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}")
@@ -50,10 +57,10 @@ class LLMProvider(LLMProviderBase):
stream=True,
tools=functions,
)
for chunk in stream:
yield chunk.choices[0].delta.content, chunk.choices[0].delta.tool_calls
except Exception as e:
logger.bind(tag=TAG).error(f"Error in Ollama function call: {e}")
yield {"type": "content", "content": f"【Ollama服务响应异常: {str(e)}"}
yield {"type": "content", "content": f"【Ollama服务响应异常: {str(e)}"}
+10 -9
View File
@@ -14,7 +14,7 @@ logger = setup_logging()
class TTSProviderBase(ABC):
def __init__(self, config, delete_audio_file):
self.delete_audio_file = delete_audio_file
self.output_file = config.get("output_file")
self.output_file = config.get("output_dir")
@abstractmethod
def generate_filename(self):
@@ -35,7 +35,7 @@ class TTSProviderBase(ABC):
return tmp_file
except Exception as e:
logger.bind(tag=TAG).info(f": {e}")
logger.bind(tag=TAG).info(f"Failed to generate TTS file: {e}")
return None
def to_tts_stream(self, text, queue: queue.Queue, text_index=0):
@@ -52,19 +52,20 @@ class TTSProviderBase(ABC):
async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0):
raise Exception("该TTS还没有实现stream模式")
def wav_to_opus_data(self, wav_file_path):
# 使用pydub加载PCM文件
def audio_to_opus_data(self, audio_file_path):
"""音频文件转换为Opus编码"""
# 获取文件后缀名
file_type = os.path.splitext(wav_file_path)[1]
file_type = os.path.splitext(audio_file_path)[1]
if file_type:
file_type = file_type.lstrip('.')
audio = AudioSegment.from_file(wav_file_path, format=file_type)
audio = AudioSegment.from_file(audio_file_path, format=file_type)
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
# 音频时长(秒)
duration = len(audio) / 1000.0
# 转换为单声道和16kHz采样率(确保与编码器匹配)
audio = audio.set_channels(1).set_frame_rate(16000)
# 获取原始PCM数据(16位小端)
raw_data = audio.raw_data
@@ -0,0 +1,35 @@
import os
import uuid
import requests
from config.logger import setup_logging
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
TAG = __name__
logger = setup_logging()
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.url = config.get("url")
self.headers = config.get("headers", {})
self.params = config.get("params")
self.format = config.get("format", "wav")
self.output_file = config.get("output_dir", "tmp/")
def generate_filename(self):
return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}.{self.format}")
async def text_to_speak(self, text, output_file):
request_params = {}
for k, v in self.params.items():
if isinstance(v, str) and "{prompt_text}" in v:
v = v.replace("{prompt_text}", text)
request_params[k] = v
resp = requests.get(self.url, params=request_params, headers=self.headers)
if resp.status_code == 200:
with open(output_file, "wb") as file:
file.write(resp.content)
else:
logger.bind(tag=TAG).error(f"Custom TTS请求失败: {resp.status_code} - {resp.text}")
@@ -177,18 +177,8 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0):
try:
# Prepare reference data
byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio]
ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text]
data = {
"text": text,
"references": [
ServeReferenceAudio(
audio=audio if audio else b"", text=text
)
for text, audio in zip(ref_texts, byte_audios)
],
"reference_id": self.reference_id,
"normalize": self.normalize,
"format": self.format,
@@ -202,6 +192,18 @@ class TTSProvider(TTSProviderBase):
"seed": self.seed,
}
# Prepare reference data
if self.reference_audio and self.reference_text:
byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio]
ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text]
data["references"] = [
ServeReferenceAudio(
audio=audio if audio else b"", text=text
)
for text, audio in zip(ref_texts, byte_audios)
],
data["reference_id"] = None
pydantic_data = ServeTTSRequest(**data)
audio_buff = None
chunk_total = b''
@@ -224,7 +226,7 @@ class TTSProvider(TTSProviderBase):
if len(chunk_total) % 2 == 0 and chunk_total[-2:] == b'\x00\x00':
audio = self._get_audio_from_tts(chunk_total)
audio_raw = audio_raw + audio.raw_data
#长度凑够2贞开始发送,60ms*4=240ms
# 长度凑够2贞开始发送,60ms*4=240ms
if len(audio_raw) >= 7680:
duration = 60 * len(audio_raw) // 1920
if (len(audio_raw) % 1920) > 0:
@@ -12,17 +12,18 @@ class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.url = config.get("url")
self.text_lang = config.get("text_lang", "audo")
self.ref_audio_path = config.get("ref_audio_path")
self.prompt_lang = config.get("prompt_lang")
self.refer_wav_path = config.get("refer_wav_path")
self.prompt_text = config.get("prompt_text")
self.top_k = config.get("top_k", 5)
self.top_p = config.get("top_p", 1)
self.temperature = config.get("temperature", 1)
self.sample_steps = config.get("sample_steps", 16)
self.media_type = config.get("media_type", "wav")
self.streaming_mode = config.get("streaming_mode", False)
self.threshold = config.get("threshold", 30)
self.prompt_language = config.get("prompt_language")
self.text_language = config.get("text_language", "audo")
self.top_k = config.get("top_k", 15)
self.top_p = config.get("top_p", 1.0)
self.temperature = config.get("temperature", 1.0)
self.cut_punc = config.get("cut_punc","")
self.speed = config.get("speed", 1.0)
self.inp_refs = config.get("inp_refs",[])
self.sample_steps = config.get("sample_steps",32)
self.if_sr = config.get("if_sr",False)
def generate_filename(self, extension=".wav"):
@@ -30,18 +31,19 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak(self, text, output_file):
request_params = {
"text": text,
"text_lang": self.text_lang,
"ref_audio_path": self.ref_audio_path,
"prompt_lang": self.prompt_lang,
"refer_wav_path": self.refer_wav_path,
"prompt_text": self.prompt_text,
"prompt_language": self.prompt_language,
"text": text,
"text_language": self.text_language,
"top_k": self.top_k,
"top_p": self.top_p,
"temperature": self.temperature,
"cut_punc": self.cut_punc,
"speed": self.speed,
"inp_refs": self.inp_refs,
"sample_steps": self.sample_steps,
"media_type": self.media_type,
"streaming_mode": self.streaming_mode,
"threshold": self.threshold,
"if_sr": self.if_sr,
}
resp = requests.get(self.url, params=request_params)
@@ -14,7 +14,7 @@ class TTSProvider(TTSProviderBase):
self.voice = config.get("voice", "alloy")
self.response_format = "wav"
self.speed = config.get("speed", 1.0)
self.output_file = config.get("output_file", "tmp/")
self.output_file = config.get("output_dir", "tmp/")
check_model_key("TTS", self.api_key)
def generate_filename(self, extension=".wav"):
@@ -17,7 +17,7 @@ class TTSProvider(TTSProviderBase):
self.volume_change_dB = config.get("volume_change_dB", 0)
self.speed_factor = config.get("speed_factor", 1)
self.stream = config.get("stream", False)
self.output_file = config.get("output_file")
self.output_file = config.get("output_dir")
self.pitch_factor = config.get("pitch_factor", 0)
self.format = config.get("format", "mp3")
self.emotion = config.get("emotion", 1)
+15 -5
View File
@@ -4,10 +4,12 @@ from datetime import datetime
class Message:
def __init__(self, role: str, content: str = None, uniq_id: str = None):
def __init__(self, role: str, content: str = None, uniq_id: str = None, tool_calls = None, tool_call_id=None):
self.uniq_id = uniq_id if uniq_id is not None else str(uuid.uuid4())
self.role = role
self.content = content
self.tool_calls = tool_calls
self.tool_call_id = tool_call_id
class Dialogue:
@@ -19,10 +21,18 @@ class Dialogue:
def put(self, message: Message):
self.dialogue.append(message)
def getMessages(self, m, dialogue):
if m.tool_calls is not None:
dialogue.append({"role": m.role, "tool_calls": m.tool_calls})
elif m.role == "tool":
dialogue.append({"role": m.role, "tool_call_id": m.tool_call_id, "content": m.content})
else:
dialogue.append({"role": m.role, "content": m.content})
def get_llm_dialogue(self) -> List[Dict[str, str]]:
dialogue = []
for m in self.dialogue:
dialogue.append({"role": m.role, "content": m.content})
self.getMessages(m, dialogue)
return dialogue
def get_llm_dialogue_with_memory(self, memory_str: str = None) -> List[Dict[str, str]]:
@@ -46,8 +56,8 @@ class Dialogue:
dialogue.append({"role": "system", "content": enhanced_system_prompt})
# 添加用户和助手的对话
for msg in self.dialogue:
if msg.role != "system": # 跳过原始的系统消息
dialogue.append({"role": msg.role, "content": msg.content})
for m in self.dialogue:
if m.role != "system": # 跳过原始的系统消息
self.getMessages(m, dialogue)
return dialogue
+59
View File
@@ -5,6 +5,7 @@ import socket
import subprocess
import logging
import re
import requests
def get_project_dir():
@@ -23,6 +24,64 @@ def get_local_ip():
except Exception as e:
return "127.0.0.1"
def is_private_ip(ip_addr):
"""
Check if an IP address is a private IP address (compatible with IPv4 and IPv6).
@param {string} ip_addr - The IP address to check.
@return {bool} True if the IP address is private, False otherwise.
"""
try:
# Validate IPv4 or IPv6 address format
if not re.match(r"^(\d{1,3}\.){3}\d{1,3}$|^([0-9a-fA-F]{1,4}:){7}[0-9a-fA-F]{1,4}$", ip_addr):
return False # Invalid IP address format
# IPv4 private address ranges
if '.' in ip_addr: # IPv4 address
ip_parts = list(map(int, ip_addr.split('.')))
if ip_parts[0] == 10:
return True # 10.0.0.0/8 range
elif ip_parts[0] == 172 and 16 <= ip_parts[1] <= 31:
return True # 172.16.0.0/12 range
elif ip_parts[0] == 192 and ip_parts[1] == 168:
return True # 192.168.0.0/16 range
elif ip_addr == '127.0.0.1':
return True # Loopback address
elif ip_parts[0] == 169 and ip_parts[1] == 254:
return True # Link-local address 169.254.0.0/16
else:
return False # Not a private IPv4 address
else: # IPv6 address
ip_addr = ip_addr.lower()
if ip_addr.startswith('fc00:') or ip_addr.startswith('fd00:'):
return True # Unique Local Addresses (FC00::/7)
elif ip_addr == '::1':
return True # Loopback address
elif ip_addr.startswith('fe80:'):
return True # Link-local unicast addresses (FE80::/10)
else:
return False # Not a private IPv6 address
except (ValueError, IndexError):
return False # IP address format error or insufficient segments
def get_ip_info(ip_addr):
try:
base_url = "https://freeipapi.com/api/json"
url = base_url if is_private_ip(ip_addr) else f"{base_url}/{ip_addr}"
resp = requests.get(url).json()
ip_info = {
"city": resp.get("cityName"),
"region": resp.get("regionName"),
"country": resp.get("countryName")
}
return ip_info
except Exception as e:
logging.error(f"Error getting client ip info: {e}")
return {}
def read_config(config_path):
with open(config_path, "r", encoding="utf-8") as file:
+3 -5
View File
@@ -2,7 +2,6 @@ import asyncio
import websockets
from config.logger import setup_logging
from core.connection import ConnectionHandler
from core.handle.musicHandler import MusicHandler
from core.utils.util import get_local_ip
from core.utils import asr, vad, llm, tts, memory, intent
@@ -13,7 +12,7 @@ class WebSocketServer:
def __init__(self, config: dict):
self.config = config
self.logger = setup_logging()
self._vad, self._asr, self._llm, self._tts, self._music, self._memory, self.intent = self._create_processing_instances()
self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = self._create_processing_instances()
self.active_connections = set() # 添加全局连接记录
def _create_processing_instances(self):
@@ -50,7 +49,6 @@ class WebSocketServer:
self.config["TTS"][self.config["selected_module"]["TTS"]],
self.config["delete_audio"]
),
MusicHandler(self.config),
memory.create_instance(memory_cls_name, memory_cfg),
intent.create_instance(
self.config["selected_module"]["Intent"]
@@ -66,7 +64,7 @@ class WebSocketServer:
host = server_config["ip"]
port = server_config["port"]
selected_module = self.config.get("selected_module")
self.logger.bind(tag=TAG).info(f"selected_module: {selected_module}")
self.logger.bind(tag=TAG).info(f"selected_module values: {', '.join(selected_module.values())}")
self.logger.bind(tag=TAG).info("Server is running at ws://{}:{}", get_local_ip(), port)
self.logger.bind(tag=TAG).info("=======上面的地址是websocket协议地址,请勿用浏览器访问=======")
@@ -80,7 +78,7 @@ class WebSocketServer:
async def _handle_connection(self, websocket):
"""处理新连接,每次创建独立的ConnectionHandler"""
# 创建ConnectionHandler时传入当前server实例
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._music, self._memory, self.intent)
handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._memory, self.intent)
self.active_connections.add(handler)
try:
await handler.handle_connection(websocket)