mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 04:43:56 +08:00
update: 同步非流式处理 10秒超时链接不复用
This commit is contained in:
@@ -15,7 +15,7 @@ from urllib import parse
|
|||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.utils import opus_encoder_utils
|
from core.utils import opus_encoder_utils, textUtils
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -122,10 +122,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.ws_url = f"wss://{self.host}/ws/v1"
|
self.ws_url = f"wss://{self.host}/ws/v1"
|
||||||
self.ws = None
|
self.ws = None
|
||||||
self._monitor_task = None
|
self._monitor_task = None
|
||||||
|
self.last_active_time = None
|
||||||
|
|
||||||
# 文本获取
|
# 文本符号
|
||||||
self.sentence_queue = queue.Queue()
|
self.sentence_queue = None
|
||||||
self.text_buffer = ""
|
|
||||||
self.sentence_end_chars = {'.', '。', '!', '!', '?', '?', '\n', "~", ";", ";", ":", ":", " ", ",", ","}
|
self.sentence_end_chars = {'.', '。', '!', '!', '?', '?', '\n', "~", ";", ";", ":", ":", " ", ",", ","}
|
||||||
|
|
||||||
# 创建Opus编码器
|
# 创建Opus编码器
|
||||||
@@ -180,8 +180,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
if self._is_token_expired():
|
if self._is_token_expired():
|
||||||
logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...")
|
logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...")
|
||||||
self._refresh_token()
|
self._refresh_token()
|
||||||
if self.ws:
|
current_time = time.time()
|
||||||
# 10秒内才可以复用,适合连续对话
|
if self.ws and current_time - self.last_active_time < 10:
|
||||||
|
# 10秒内才可以复用链接进行连续对话
|
||||||
logger.bind(tag=TAG).info(f"使用已有链接...")
|
logger.bind(tag=TAG).info(f"使用已有链接...")
|
||||||
return self.ws
|
return self.ws
|
||||||
logger.bind(tag=TAG).info("开始建立新连接...")
|
logger.bind(tag=TAG).info("开始建立新连接...")
|
||||||
@@ -194,10 +195,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
close_timeout=10,
|
close_timeout=10,
|
||||||
)
|
)
|
||||||
logger.bind(tag=TAG).info("WebSocket连接建立成功")
|
logger.bind(tag=TAG).info("WebSocket连接建立成功")
|
||||||
|
self.last_active_time = time.time()
|
||||||
return self.ws
|
return self.ws
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"建立连接失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"建立连接失败: {str(e)}")
|
||||||
self.ws = None
|
self.ws = None
|
||||||
|
self.last_active_time = None
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def tts_text_priority_thread(self):
|
def tts_text_priority_thread(self):
|
||||||
@@ -223,8 +226,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.conn.sentence_id = uuid.uuid4().hex
|
self.conn.sentence_id = uuid.uuid4().hex
|
||||||
logger.bind(tag=TAG).info(f"自动生成新的 会话ID: {self.conn.sentence_id}")
|
logger.bind(tag=TAG).info(f"自动生成新的 会话ID: {self.conn.sentence_id}")
|
||||||
|
|
||||||
# aliyun独有的message_id需要自己生成
|
# aliyunStream独有的参数生成
|
||||||
self.conn.message_id = str(uuid.uuid4().hex)
|
self.conn.message_id = str(uuid.uuid4().hex)
|
||||||
|
self.sentence_queue = queue.Queue()
|
||||||
|
self.text_buffer = ""
|
||||||
|
|
||||||
logger.bind(tag=TAG).info("开始启动TTS会话...")
|
logger.bind(tag=TAG).info("开始启动TTS会话...")
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
@@ -247,7 +252,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
self.text_buffer += message.content_detail
|
self.text_buffer += message.content_detail
|
||||||
if message.content_detail in self.sentence_end_chars and len(self.text_buffer) > 6:
|
if message.content_detail in self.sentence_end_chars and len(self.text_buffer) > 6:
|
||||||
self.sentence_queue.put(self.text_buffer)
|
self.sentence_queue.put(textUtils.get_string_no_punctuation_or_emoji(self.text_buffer))
|
||||||
self.text_buffer = ""
|
self.text_buffer = ""
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
self.text_to_speak(message.content_detail, None),
|
self.text_to_speak(message.content_detail, None),
|
||||||
@@ -273,8 +278,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
||||||
self.sentence_queue.put(self.text_buffer)
|
self.sentence_queue.put(textUtils.get_string_no_punctuation_or_emoji(self.text_buffer))
|
||||||
print(list(self.sentence_queue.queue))
|
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
self.finish_session(self.conn.sentence_id),
|
self.finish_session(self.conn.sentence_id),
|
||||||
loop=self.conn.loop,
|
loop=self.conn.loop,
|
||||||
@@ -310,6 +314,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
await self.ws.send(json.dumps(run_request))
|
await self.ws.send(json.dumps(run_request))
|
||||||
|
self.last_active_time = time.time()
|
||||||
return
|
return
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -359,6 +364,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
await self.ws.send(json.dumps(start_request))
|
await self.ws.send(json.dumps(start_request))
|
||||||
|
self.last_active_time = time.time()
|
||||||
logger.bind(tag=TAG).info("会话启动请求已发送")
|
logger.bind(tag=TAG).info("会话启动请求已发送")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
|
||||||
@@ -381,6 +387,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
}
|
}
|
||||||
await self.ws.send(json.dumps(stop_request))
|
await self.ws.send(json.dumps(stop_request))
|
||||||
logger.bind(tag=TAG).info("会话结束请求已发送")
|
logger.bind(tag=TAG).info("会话结束请求已发送")
|
||||||
|
self.last_active_time = time.time()
|
||||||
if self._monitor_task:
|
if self._monitor_task:
|
||||||
try:
|
try:
|
||||||
await self._monitor_task
|
await self._monitor_task
|
||||||
@@ -414,6 +421,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
except:
|
except:
|
||||||
pass
|
pass
|
||||||
self.ws = None
|
self.ws = None
|
||||||
|
self.last_active_time = None
|
||||||
|
|
||||||
async def _start_monitor_tts_response(self):
|
async def _start_monitor_tts_response(self):
|
||||||
"""监听TTS响应"""
|
"""监听TTS响应"""
|
||||||
@@ -426,6 +434,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
try:
|
try:
|
||||||
msg = await self.ws.recv()
|
msg = await self.ws.recv()
|
||||||
|
self.last_active_time = time.time()
|
||||||
# 检查客户端是否中止
|
# 检查客户端是否中止
|
||||||
if self.conn.client_abort:
|
if self.conn.client_abort:
|
||||||
logger.bind(tag=TAG).info("收到打断信息,终止监听TTS响应")
|
logger.bind(tag=TAG).info("收到打断信息,终止监听TTS响应")
|
||||||
@@ -459,6 +468,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
logger.bind(tag=TAG).debug(f"会话结束~~")
|
logger.bind(tag=TAG).debug(f"会话结束~~")
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
session_finished = True
|
session_finished = True
|
||||||
|
self.reuse_judgment = time.time()
|
||||||
break
|
break
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
logger.bind(tag=TAG).warning("收到无效的JSON消息")
|
logger.bind(tag=TAG).warning("收到无效的JSON消息")
|
||||||
@@ -502,65 +512,102 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
def to_tts(self, text: str) -> list:
|
||||||
"""非流式TTS处理,用于测试及保存音频文件的场景"""
|
"""非流式TTS处理,用于测试及保存音频文件的场景"""
|
||||||
start_time = time.time()
|
|
||||||
text = MarkdownCleaner.clean_markdown(text)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# 使用同步方式进行TTS转换
|
# 创建新的事件循环
|
||||||
if self._is_token_expired():
|
loop = asyncio.new_event_loop()
|
||||||
self._refresh_token()
|
asyncio.set_event_loop(loop)
|
||||||
|
|
||||||
# 构造请求数据
|
# 生成会话ID
|
||||||
request_json = {
|
session_id = uuid.uuid4().hex
|
||||||
"appkey": self.appkey,
|
message_id = uuid.uuid4().hex
|
||||||
"token": self.token,
|
# 存储音频数据
|
||||||
"text": text,
|
audio_data = []
|
||||||
"format": "pcm",
|
|
||||||
"sample_rate": self.sample_rate,
|
|
||||||
"voice": self.voice,
|
|
||||||
"volume": self.volume,
|
|
||||||
"speech_rate": self.speech_rate,
|
|
||||||
"pitch_rate": self.pitch_rate,
|
|
||||||
}
|
|
||||||
|
|
||||||
# 使用HTTP接口进行同步请求
|
async def _generate_audio():
|
||||||
import requests
|
# 刷新Token(如果需要)
|
||||||
api_url = f"https://{self.host}/stream/v1/tts"
|
if self._is_token_expired():
|
||||||
headers = {"Content-Type": "application/json"}
|
self._refresh_token()
|
||||||
|
|
||||||
resp = requests.post(api_url, json=request_json, headers=headers)
|
# 建立WebSocket连接
|
||||||
|
ws = await websockets.connect(
|
||||||
if resp.status_code == 401: # Token过期特殊处理
|
self.ws_url,
|
||||||
self._refresh_token()
|
additional_headers={"X-NLS-Token": self.token},
|
||||||
resp = requests.post(api_url, json=request_json, headers=headers)
|
ping_interval=30,
|
||||||
|
ping_timeout=10,
|
||||||
if resp.headers["Content-Type"].startswith("audio/"):
|
close_timeout=10,
|
||||||
pcm_data = resp.content
|
|
||||||
|
|
||||||
# 使用opus编码器处理PCM数据
|
|
||||||
opus_datas = []
|
|
||||||
frame_bytes = int(
|
|
||||||
self.opus_encoder.sample_rate
|
|
||||||
* self.opus_encoder.channels
|
|
||||||
* self.opus_encoder.frame_size_ms
|
|
||||||
/ 1000
|
|
||||||
* 2
|
|
||||||
)
|
)
|
||||||
|
try:
|
||||||
|
# 发送StartSynthesis请求
|
||||||
|
start_request = {
|
||||||
|
"header": {
|
||||||
|
"message_id": message_id,
|
||||||
|
"task_id": session_id,
|
||||||
|
"namespace": "FlowingSpeechSynthesizer",
|
||||||
|
"name": "StartSynthesis",
|
||||||
|
"appkey": self.appkey,
|
||||||
|
},
|
||||||
|
"payload": {
|
||||||
|
"voice": self.voice,
|
||||||
|
"format": self.format,
|
||||||
|
"sample_rate": self.sample_rate,
|
||||||
|
"volume": self.volume,
|
||||||
|
"speech_rate": self.speech_rate,
|
||||||
|
"pitch_rate": self.pitch_rate,
|
||||||
|
"enable_subtitle": True
|
||||||
|
}
|
||||||
|
}
|
||||||
|
await ws.send(json.dumps(start_request))
|
||||||
|
|
||||||
# 分帧处理PCM数据
|
# 发送文本合成请求
|
||||||
for i in range(0, len(pcm_data), frame_bytes):
|
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||||
frame = pcm_data[i:i + frame_bytes]
|
run_request = {
|
||||||
if len(frame) == frame_bytes:
|
"header": {
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(frame, False)
|
"message_id": message_id,
|
||||||
if opus:
|
"task_id": session_id,
|
||||||
opus_datas.extend(opus)
|
"namespace": "FlowingSpeechSynthesizer",
|
||||||
|
"name": "RunSynthesis",
|
||||||
|
"appkey": self.appkey,
|
||||||
|
},
|
||||||
|
"payload": {
|
||||||
|
"text": filtered_text
|
||||||
|
}
|
||||||
|
}
|
||||||
|
await ws.send(json.dumps(run_request))
|
||||||
|
|
||||||
logger.bind(tag=TAG).info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒")
|
# 发送停止合成请求
|
||||||
return opus_datas
|
stop_request = {
|
||||||
else:
|
"header": {
|
||||||
logger.bind(tag=TAG).error(f"TTS请求失败: {resp.content}")
|
"message_id": message_id,
|
||||||
return []
|
"task_id": session_id,
|
||||||
|
"namespace": "FlowingSpeechSynthesizer",
|
||||||
|
"name": "StopSynthesis",
|
||||||
|
"appkey": self.appkey,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
await ws.send(json.dumps(stop_request))
|
||||||
|
|
||||||
|
# 接收音频数据
|
||||||
|
while True:
|
||||||
|
msg = await ws.recv()
|
||||||
|
if isinstance(msg, (bytes, bytearray)):
|
||||||
|
# 编码为Opus并收集
|
||||||
|
opus_frames = self.opus_encoder.encode_pcm_to_opus(msg, False)
|
||||||
|
audio_data.extend(opus_frames)
|
||||||
|
elif isinstance(msg, str):
|
||||||
|
data = json.loads(msg)
|
||||||
|
header = data.get("header", {})
|
||||||
|
if header.get("name") == "SynthesisCompleted":
|
||||||
|
break
|
||||||
|
finally:
|
||||||
|
try:
|
||||||
|
await ws.close()
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
|
||||||
|
loop.run_until_complete(_generate_audio())
|
||||||
|
loop.close()
|
||||||
|
|
||||||
|
return audio_data
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
|
||||||
return []
|
return []
|
||||||
|
|||||||
Reference in New Issue
Block a user