update:优化

This commit is contained in:
hrz
2025-05-25 08:56:58 +08:00
parent 7a598d5839
commit 40632019ac
4 changed files with 28 additions and 18 deletions
+4
View File
@@ -123,6 +123,7 @@ class ConnectionHandler:
# tts相关变量
self.tts_first_text_index = -1
self.tts_last_text_index = -1
self.tts_session_id = None
# iot相关变量
self.iot_descriptors = {}
@@ -521,6 +522,9 @@ class ConnectionHandler:
)
memory_str = future.result()
uuid_str = str(uuid.uuid4()).replace("-", "")
self.tts_session_id = uuid_str
if functions is not None:
# 使用支持functions的streaming接口
llm_responses = self.llm.response_with_functions(
@@ -58,6 +58,7 @@ async def sendAudioMessage(conn, audios, text, text_index=0):
# 发送结束消息(如果是最后一个文本)
if conn.llm_finish_task and text_index == conn.tts_last_text_index:
await send_tts_message(conn, "stop", None)
await conn.tts.finish_session(conn.tts_session_id)
if conn.close_after_chat:
await conn.close()
@@ -40,17 +40,6 @@ class TTSProviderBase(ABC):
pass
def to_tts(self, text, index):
"""如果是流式实现,一般没有文件生成,我们返回枚举值"""
if self.interface_type != TTSImplementationType.NON_STREAMING:
if index == 1:
future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.session_id), loop=self.conn.loop
)
future.result()
asyncio.run(self.text_to_speak(text, None))
return self.interface_type.value
"""以下是非流式实现,会返回文件"""
tmp_file = self.generate_filename()
try:
max_repeat_time = 5
@@ -173,6 +173,17 @@ class TTSProvider(TTSProviderBase):
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
)
def to_tts(self, text, index):
if index == self.conn.tts_first_text_index:
future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.tts_session_id), loop=self.conn.loop
)
future.result()
future = asyncio.run_coroutine_threadsafe(
self.text_to_speak(text, None), loop=self.conn.loop
)
future.result()
return self.interface_type.value
async def send_event(
self, header: bytes, optional: bytes | None = None, payload: bytes = None
):
@@ -327,6 +338,7 @@ class TTSProvider(TTSProviderBase):
optional = Optional(event=EVENT_StartSession, sessionId=session_id).as_bytes()
payload = self.get_payload_bytes(event=EVENT_StartSession, speaker=self.speaker)
await self.send_event(header, optional, payload)
logger.bind(tag=TAG).info(f"会话开始~~{session_id}")
async def finish_session(self, session_id):
self.stop_event_response.set()
@@ -357,7 +369,8 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak(self, text, _):
# 发送文本
await self.send_text(self.speaker, text, self.conn.session_id)
await self.send_text(self.speaker, text, self.conn.tts_session_id)
logger.bind(tag=TAG).info(f"发送文本~~{text}")
return
def _start_monitor_tts_response_thread(self):
@@ -383,21 +396,24 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).info(
f"推送数据到队列里面帧数~~{len(opus_datas)}"
)
self.audio_play_queue.put((opus_datas, None, 0))
self.audio_play_queue.put(
(opus_datas, None, self.conn.tts_last_text_index - 1)
)
elif res.optional.event == EVENT_TTSSentenceStart:
json_data = json.loads(res.payload.decode("utf-8"))
self.tts_text = json_data.get("text", "")
logger.bind(tag=TAG).info(f"句子开始~~{self.tts_text}")
self.audio_play_queue.put(([], self.tts_text, 0))
self.audio_play_queue.put(
([], self.tts_text, self.conn.tts_first_text_index)
)
elif res.optional.event == EVENT_TTSSentenceEnd:
logger.bind(tag=TAG).info(f"句子结束~~{self.tts_text}")
self.audio_play_queue.put(([], self.tts_text, 0))
self.audio_play_queue.put(
([], self.tts_text, self.conn.tts_last_text_index)
)
elif res.optional.event == EVENT_SessionFinished:
logger.bind(tag=TAG).info(f"会话结束~~,最后一句补零")
opus_datas = pcm_to_data(b"")
self.audio_play_queue.put((opus_datas, self.tts_text, 0))
else:
continue
except websockets.ConnectionClosed:
break # 连接关闭时退出监听