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 sendAudio(conn, audios, False)
await sendAudio(conn, audios, pre_buffer)
await send_tts_message(conn, "sentence_end", text)
# 发送结束消息(如果是最后一个文本)
if conn.llm_finish_task and sentenceType == SentenceType.LAST:
await send_tts_message(conn, "stop", None)
await conn.tts.finish_session(conn.sentence_id)
if conn.close_after_chat:
await conn.close()
@@ -1,4 +1,3 @@
import os
import uuid
import json
import hmac
+11 -7
View File
@@ -1,17 +1,19 @@
import asyncio
from config.logger import setup_logging
import queue
import os
import queue
import uuid
import datetime
import asyncio
import threading
from core.utils import p3
from datetime import datetime
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.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
TAG = __name__
@@ -26,6 +28,7 @@ class TTSProviderBase(ABC):
self.output_file = config.get("output_dir")
self.tts_text_queue = queue.Queue()
self.tts_audio_queue = queue.Queue()
self.tts_audio_first_sentence = True
self.tts_text_buff = []
self.punctuations = (
@@ -173,6 +176,7 @@ class TTSProviderBase(ABC):
self.processed_chars = 0
self.tts_text_buff = []
self.is_first_sentence = True
self.tts_audio_first_sentence = True
elif ContentType.TEXT == message.content_type:
self.tts_text_buff.append(message.content_detail)
segment_text = self._get_segment_text()
@@ -1,9 +1,4 @@
import os
import uuid
import json
import base64
import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
@@ -1,5 +1,4 @@
import os
import asyncio
from config.logger import setup_logging
from core.providers.tts.base import TTSProviderBase
@@ -1,9 +1,7 @@
import os
import uuid
import json
import base64
import requests
from datetime import datetime
from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase
from config.logger import setup_logging
@@ -1,12 +1,9 @@
import base64
import os
import uuid
import requests
import ormsgpack
from pathlib import Path
from pydantic import BaseModel, Field, conint, model_validator
from typing_extensions import Annotated
from datetime import datetime
from typing import Literal
from core.utils.util import check_model_key, parse_string_to_list
from core.providers.tts.base import TTSProviderBase
@@ -1,10 +1,5 @@
import os
import uuid
import json
import base64
import requests
from config.logger import setup_logging
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
from core.utils.util import parse_string_to_list
@@ -1,8 +1,5 @@
import os
import uuid
import requests
from config.logger import setup_logging
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
from core.utils.util import parse_string_to_list
@@ -7,7 +7,7 @@ import websockets
import queue
from config.logger import setup_logging
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
TAG = __name__
@@ -148,6 +148,8 @@ class TTSProvider(TTSProviderBase):
self.enable_two_way = True
self.start_connection_flag = False
self.tts_text = ""
# 合成文字语音后,播放的音频文件列表
self.tts_audio_files = []
###################################################################################
# 火山双流式TTS重写父类的方法--开始
@@ -173,12 +175,16 @@ class TTSProvider(TTSProviderBase):
while not self.conn.stop_event.is_set():
try:
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:
# 初始化参数
future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.sentence_id), loop=self.conn.loop
)
future.result()
self.tts_audio_first_sentence = True
elif ContentType.TEXT == message.content_type:
if message.content_detail:
future = asyncio.run_coroutine_threadsafe(
@@ -205,7 +211,6 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak(self, text, _):
# 发送文本
await self.send_text(self.speaker, text, self.conn.sentence_id)
logger.bind(tag=TAG).info(f"发送文本~~{text}")
return
###################################################################################
@@ -227,22 +232,22 @@ class TTSProvider(TTSProviderBase):
if 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}")
logger.bind(tag=TAG).info(f"语音生成成功: {self.tts_text}")
self.tts_audio_queue.put((SentenceType.FIRST, [], self.tts_text))
elif (
res.optional.event == EVENT_TTSResponse
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)
logger.bind(tag=TAG).info(
logger.bind(tag=TAG).debug(
f"推送数据到队列里面帧数~~{len(opus_datas)}"
)
self.tts_audio_queue.put((SentenceType.MIDDLE, opus_datas, None))
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:
logger.bind(tag=TAG).info(f"会话结束~~")
logger.bind(tag=TAG).debug(f"会话结束~~")
self.tts_audio_queue.put((SentenceType.LAST, [], None))
continue
except websockets.ConnectionClosed:
@@ -355,8 +360,8 @@ class TTSProvider(TTSProviderBase):
return await self.send_event(header, optional, payload)
def print_response(self, res, tag_msg: str):
logger.bind(tag=TAG).info(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} header:{res.header.__dict__}")
logger.bind(tag=TAG).debug(f"===>{tag_msg} optional:{res.optional.__dict__}")
def get_payload_bytes(
self,
@@ -405,10 +410,10 @@ 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}")
logger.bind(tag=TAG).debug(f"开始会话~~{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(
message_type=FULL_CLIENT_REQUEST,
message_type_specific_flags=MsgTypeFlagWithEvent,
@@ -1,7 +1,4 @@
import os
import uuid
import requests
from datetime import datetime
from core.utils.util import check_model_key
from core.providers.tts.base import TTSProviderBase
from config.logger import setup_logging
@@ -1,7 +1,4 @@
import os
import uuid
import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
@@ -1,6 +1,5 @@
import hashlib
import hmac
import os
import time
import uuid
import json
@@ -82,4 +82,4 @@ class TTSProvider(TTSProviderBase):
except Exception as e:
print("error:", e)
raise Exception(f"{__name__}: TTS请求失败")
raise Exception(f"{__name__}: TTS请求失败")