mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-27 09:33:55 +08:00
mergin main,速度提升一下
This commit is contained in:
@@ -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 ""
|
||||
@@ -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)}】"}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user