update:优化代码

This commit is contained in:
hrz
2025-05-26 16:12:38 +08:00
parent ae64233986
commit 9787ca60da
14 changed files with 34 additions and 48 deletions
@@ -47,17 +47,21 @@ async def sendAudioMessage(conn, sentenceType, audios, text):
} }
) )
) )
pre_buffer = False
if conn.tts.tts_audio_first_sentence and text is not None:
conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
conn.tts.tts_audio_first_sentence = False
pre_buffer = True
await send_tts_message(conn, "sentence_start", text) await send_tts_message(conn, "sentence_start", text)
await sendAudio(conn, audios, False) await sendAudio(conn, audios, pre_buffer)
await send_tts_message(conn, "sentence_end", text) await send_tts_message(conn, "sentence_end", text)
# 发送结束消息(如果是最后一个文本) # 发送结束消息(如果是最后一个文本)
if conn.llm_finish_task and sentenceType == SentenceType.LAST: if conn.llm_finish_task and sentenceType == SentenceType.LAST:
await send_tts_message(conn, "stop", None) await send_tts_message(conn, "stop", None)
await conn.tts.finish_session(conn.sentence_id)
if conn.close_after_chat: if conn.close_after_chat:
await conn.close() await conn.close()
@@ -1,4 +1,3 @@
import os
import uuid import uuid
import json import json
import hmac import hmac
+11 -7
View File
@@ -1,17 +1,19 @@
import asyncio
from config.logger import setup_logging
import queue
import os import os
import queue
import uuid import uuid
import datetime import asyncio
import threading import threading
from core.utils import p3 from core.utils import p3
from datetime import datetime
from core.utils import textUtils from core.utils import textUtils
from abc import ABC, abstractmethod
from config.logger import setup_logging
from core.utils.util import audio_to_data
from core.utils.tts import MarkdownCleaner
from core.handle.sendAudioHandle import sendAudioMessage from core.handle.sendAudioHandle import sendAudioMessage
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType, ContentType from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType, ContentType
from abc import ABC, abstractmethod
from core.utils.tts import MarkdownCleaner
from core.utils.util import audio_to_data
import traceback import traceback
TAG = __name__ TAG = __name__
@@ -26,6 +28,7 @@ class TTSProviderBase(ABC):
self.output_file = config.get("output_dir") self.output_file = config.get("output_dir")
self.tts_text_queue = queue.Queue() self.tts_text_queue = queue.Queue()
self.tts_audio_queue = queue.Queue() self.tts_audio_queue = queue.Queue()
self.tts_audio_first_sentence = True
self.tts_text_buff = [] self.tts_text_buff = []
self.punctuations = ( self.punctuations = (
@@ -173,6 +176,7 @@ class TTSProviderBase(ABC):
self.processed_chars = 0 self.processed_chars = 0
self.tts_text_buff = [] self.tts_text_buff = []
self.is_first_sentence = True self.is_first_sentence = True
self.tts_audio_first_sentence = True
elif ContentType.TEXT == message.content_type: elif ContentType.TEXT == message.content_type:
self.tts_text_buff.append(message.content_detail) self.tts_text_buff.append(message.content_detail)
segment_text = self._get_segment_text() segment_text = self._get_segment_text()
@@ -1,9 +1,4 @@
import os
import uuid
import json
import base64
import requests import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
@@ -1,5 +1,4 @@
import os import os
import asyncio
from config.logger import setup_logging from config.logger import setup_logging
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
@@ -1,9 +1,7 @@
import os
import uuid import uuid
import json import json
import base64 import base64
import requests import requests
from datetime import datetime
from core.utils.util import check_model_key from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
from config.logger import setup_logging from config.logger import setup_logging
@@ -1,12 +1,9 @@
import base64 import base64
import os
import uuid
import requests import requests
import ormsgpack import ormsgpack
from pathlib import Path from pathlib import Path
from pydantic import BaseModel, Field, conint, model_validator from pydantic import BaseModel, Field, conint, model_validator
from typing_extensions import Annotated from typing_extensions import Annotated
from datetime import datetime
from typing import Literal from typing import Literal
from core.utils.util import check_model_key, parse_string_to_list from core.utils.util import check_model_key, parse_string_to_list
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
@@ -1,10 +1,5 @@
import os
import uuid
import json
import base64
import requests import requests
from config.logger import setup_logging from config.logger import setup_logging
from datetime import datetime
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
from core.utils.util import parse_string_to_list from core.utils.util import parse_string_to_list
@@ -1,8 +1,5 @@
import os
import uuid
import requests import requests
from config.logger import setup_logging from config.logger import setup_logging
from datetime import datetime
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
from core.utils.util import parse_string_to_list from core.utils.util import parse_string_to_list
@@ -7,7 +7,7 @@ import websockets
import queue import queue
from config.logger import setup_logging from config.logger import setup_logging
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType, ContentType from core.providers.tts.dto.dto import SentenceType, ContentType
from core.utils.util import pcm_to_data from core.utils.util import pcm_to_data
TAG = __name__ TAG = __name__
@@ -148,6 +148,8 @@ class TTSProvider(TTSProviderBase):
self.enable_two_way = True self.enable_two_way = True
self.start_connection_flag = False self.start_connection_flag = False
self.tts_text = "" self.tts_text = ""
# 合成文字语音后,播放的音频文件列表
self.tts_audio_files = []
################################################################################### ###################################################################################
# 火山双流式TTS重写父类的方法--开始 # 火山双流式TTS重写父类的方法--开始
@@ -173,12 +175,16 @@ class TTSProvider(TTSProviderBase):
while not self.conn.stop_event.is_set(): while not self.conn.stop_event.is_set():
try: try:
message = self.tts_text_queue.get(timeout=1) message = self.tts_text_queue.get(timeout=1)
logger.bind(tag=TAG).debug(
f"TTS任务|{message.sentence_type.name} {message.content_type.name}"
)
if message.sentence_type == SentenceType.FIRST: if message.sentence_type == SentenceType.FIRST:
# 初始化参数 # 初始化参数
future = asyncio.run_coroutine_threadsafe( future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.sentence_id), loop=self.conn.loop self.start_session(self.conn.sentence_id), loop=self.conn.loop
) )
future.result() future.result()
self.tts_audio_first_sentence = True
elif ContentType.TEXT == message.content_type: elif ContentType.TEXT == message.content_type:
if message.content_detail: if message.content_detail:
future = asyncio.run_coroutine_threadsafe( future = asyncio.run_coroutine_threadsafe(
@@ -205,7 +211,6 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak(self, text, _): async def text_to_speak(self, text, _):
# 发送文本 # 发送文本
await self.send_text(self.speaker, text, self.conn.sentence_id) await self.send_text(self.speaker, text, self.conn.sentence_id)
logger.bind(tag=TAG).info(f"发送文本~~{text}")
return return
################################################################################### ###################################################################################
@@ -227,22 +232,22 @@ class TTSProvider(TTSProviderBase):
if res.optional.event == EVENT_TTSSentenceStart: if res.optional.event == EVENT_TTSSentenceStart:
json_data = json.loads(res.payload.decode("utf-8")) json_data = json.loads(res.payload.decode("utf-8"))
self.tts_text = json_data.get("text", "") self.tts_text = json_data.get("text", "")
logger.bind(tag=TAG).info(f"句子开始~~{self.tts_text}") logger.bind(tag=TAG).info(f"语音生成成功: {self.tts_text}")
self.tts_audio_queue.put((SentenceType.FIRST, [], self.tts_text)) self.tts_audio_queue.put((SentenceType.FIRST, [], self.tts_text))
elif ( elif (
res.optional.event == EVENT_TTSResponse res.optional.event == EVENT_TTSResponse
and res.header.message_type == AUDIO_ONLY_RESPONSE and res.header.message_type == AUDIO_ONLY_RESPONSE
): ):
logger.bind(tag=TAG).info(f"推送数据到队列里面~~") logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
opus_datas = pcm_to_data(res.payload) opus_datas = pcm_to_data(res.payload)
logger.bind(tag=TAG).info( logger.bind(tag=TAG).debug(
f"推送数据到队列里面帧数~~{len(opus_datas)}" f"推送数据到队列里面帧数~~{len(opus_datas)}"
) )
self.tts_audio_queue.put((SentenceType.MIDDLE, opus_datas, None)) self.tts_audio_queue.put((SentenceType.MIDDLE, opus_datas, None))
elif res.optional.event == EVENT_TTSSentenceEnd: elif res.optional.event == EVENT_TTSSentenceEnd:
logger.bind(tag=TAG).info(f"句子结束~~{self.tts_text}") logger.bind(tag=TAG).debug(f"句子结束~~{self.tts_text}")
elif res.optional.event == EVENT_SessionFinished: elif res.optional.event == EVENT_SessionFinished:
logger.bind(tag=TAG).info(f"会话结束~~") logger.bind(tag=TAG).debug(f"会话结束~~")
self.tts_audio_queue.put((SentenceType.LAST, [], None)) self.tts_audio_queue.put((SentenceType.LAST, [], None))
continue continue
except websockets.ConnectionClosed: except websockets.ConnectionClosed:
@@ -355,8 +360,8 @@ class TTSProvider(TTSProviderBase):
return await self.send_event(header, optional, payload) return await self.send_event(header, optional, payload)
def print_response(self, res, tag_msg: str): def print_response(self, res, tag_msg: str):
logger.bind(tag=TAG).info(f"===>{tag_msg} header:{res.header.__dict__}") logger.bind(tag=TAG).debug(f"===>{tag_msg} header:{res.header.__dict__}")
logger.bind(tag=TAG).info(f"===>{tag_msg} optional:{res.optional.__dict__}") logger.bind(tag=TAG).debug(f"===>{tag_msg} optional:{res.optional.__dict__}")
def get_payload_bytes( def get_payload_bytes(
self, self,
@@ -405,10 +410,10 @@ class TTSProvider(TTSProviderBase):
optional = Optional(event=EVENT_StartSession, sessionId=session_id).as_bytes() optional = Optional(event=EVENT_StartSession, sessionId=session_id).as_bytes()
payload = self.get_payload_bytes(event=EVENT_StartSession, speaker=self.speaker) payload = self.get_payload_bytes(event=EVENT_StartSession, speaker=self.speaker)
await self.send_event(header, optional, payload) await self.send_event(header, optional, payload)
logger.bind(tag=TAG).info(f"会话开始~~{session_id}") logger.bind(tag=TAG).debug(f"开始会话~~{session_id}")
async def finish_session(self, session_id): async def finish_session(self, session_id):
logger.bind(tag=TAG).info(f"会话结束~~{session_id}") logger.bind(tag=TAG).debug(f"关闭会话~~{session_id}")
header = Header( header = Header(
message_type=FULL_CLIENT_REQUEST, message_type=FULL_CLIENT_REQUEST,
message_type_specific_flags=MsgTypeFlagWithEvent, message_type_specific_flags=MsgTypeFlagWithEvent,
@@ -1,7 +1,4 @@
import os
import uuid
import requests import requests
from datetime import datetime
from core.utils.util import check_model_key from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
from config.logger import setup_logging from config.logger import setup_logging
@@ -1,7 +1,4 @@
import os
import uuid
import requests import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase from core.providers.tts.base import TTSProviderBase
@@ -1,6 +1,5 @@
import hashlib import hashlib
import hmac import hmac
import os
import time import time
import uuid import uuid
import json import json
@@ -82,4 +82,4 @@ class TTSProvider(TTSProviderBase):
except Exception as e: except Exception as e:
print("error:", e) print("error:", e)
raise Exception(f"{__name__}: TTS请求失败") raise Exception(f"{__name__}: TTS请求失败")