mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 02:43:55 +08:00
update:优化代码
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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请求失败")
|
||||||
|
|||||||
Reference in New Issue
Block a user