update:修复cosyvoice非流式合成方法

This commit is contained in:
hrz
2025-07-22 13:25:54 +08:00
parent 6740d6a723
commit 60cbe1571c
@@ -72,6 +72,7 @@ class AccessToken:
) )
import requests import requests
response = requests.get(full_url) response = requests.get(full_url)
if response.ok: if response.ok:
root_obj = response.json() root_obj = response.json()
@@ -126,7 +127,7 @@ class TTSProvider(TTSProviderBase):
# 专属tts设置 # 专属tts设置
self.message_id = "" self.message_id = ""
self.tts_text = '' self.tts_text = ""
self.text_buffer = [] self.text_buffer = []
# 创建Opus编码器 # 创建Opus编码器
@@ -222,7 +223,9 @@ class TTSProvider(TTSProviderBase):
try: try:
if not getattr(self.conn, "sentence_id", None): if not getattr(self.conn, "sentence_id", None):
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}"
)
# aliyunStream独有的参数生成 # aliyunStream独有的参数生成
self.message_id = str(uuid.uuid4().hex) self.message_id = str(uuid.uuid4().hex)
@@ -273,7 +276,7 @@ class TTSProvider(TTSProviderBase):
try: try:
logger.bind(tag=TAG).info("开始结束TTS会话...") logger.bind(tag=TAG).info("开始结束TTS会话...")
self.tts_text = textUtils.get_string_no_punctuation_or_emoji( self.tts_text = textUtils.get_string_no_punctuation_or_emoji(
''.join(self.text_buffer).replace('\n', '') "".join(self.text_buffer).replace("\n", "")
) )
future = asyncio.run_coroutine_threadsafe( future = asyncio.run_coroutine_threadsafe(
self.finish_session(self.conn.sentence_id), self.finish_session(self.conn.sentence_id),
@@ -305,9 +308,7 @@ class TTSProvider(TTSProviderBase):
"name": "RunSynthesis", "name": "RunSynthesis",
"appkey": self.appkey, "appkey": self.appkey,
}, },
"payload": { "payload": {"text": filtered_text},
"text": filtered_text
}
} }
await self.ws.send(json.dumps(run_request)) await self.ws.send(json.dumps(run_request))
self.last_active_time = time.time() self.last_active_time = time.time()
@@ -327,12 +328,14 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).info(f"开始会话~~{session_id}") logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
try: try:
# 会话开始时检测上个会话的监听状态 # 会话开始时检测上个会话的监听状态
if( if (
self._monitor_task is not None self._monitor_task is not None
and isinstance(self._monitor_task, Task) and isinstance(self._monitor_task, Task)
and not self._monitor_task.done() and not self._monitor_task.done()
): ):
logger.bind(tag=TAG).info("检测到未完成的上个会话,关闭监听任务和连接...") logger.bind(tag=TAG).info(
"检测到未完成的上个会话,关闭监听任务和连接..."
)
await self.close() await self.close()
# 建立新连接 # 建立新连接
@@ -356,8 +359,8 @@ class TTSProvider(TTSProviderBase):
"volume": self.volume, "volume": self.volume,
"speech_rate": self.speech_rate, "speech_rate": self.speech_rate,
"pitch_rate": self.pitch_rate, "pitch_rate": self.pitch_rate,
"enable_subtitle": True "enable_subtitle": True,
} },
} }
await self.ws.send(json.dumps(start_request)) await self.ws.send(json.dumps(start_request))
self.last_active_time = time.time() self.last_active_time = time.time()
@@ -442,12 +445,21 @@ class TTSProvider(TTSProviderBase):
if event_name == "SynthesisStarted": if event_name == "SynthesisStarted":
logger.bind(tag=TAG).debug("TTS合成已启动") logger.bind(tag=TAG).debug("TTS合成已启动")
elif event_name == "SentenceBegin": elif event_name == "SentenceBegin":
logger.bind(tag=TAG).debug(f"句子语音生成开始: {self.tts_text}") logger.bind(tag=TAG).debug(
f"句子语音生成开始: {self.tts_text}"
)
opus_datas_cache = [] opus_datas_cache = []
self.tts_audio_queue.put((SentenceType.FIRST, [], self.tts_text)) self.tts_audio_queue.put(
(SentenceType.FIRST, [], self.tts_text)
)
elif event_name == "SentenceEnd": elif event_name == "SentenceEnd":
logger.bind(tag=TAG).info(f"句子语音生成成功: {self.tts_text}") logger.bind(tag=TAG).info(
if not is_first_sentence or first_sentence_segment_count > 10: f"句子语音生成成功: {self.tts_text}"
)
if (
not is_first_sentence
or first_sentence_segment_count > 10
):
# 发送缓存的数据 # 发送缓存的数据
self.tts_audio_queue.put( self.tts_audio_queue.put(
(SentenceType.MIDDLE, opus_datas_cache, None) (SentenceType.MIDDLE, opus_datas_cache, None)
@@ -459,7 +471,7 @@ class TTSProvider(TTSProviderBase):
self._process_before_stop_play_files() self._process_before_stop_play_files()
session_finished = True session_finished = True
self.reuse_judgment = time.time() self.reuse_judgment = time.time()
self.tts_text = '' self.tts_text = ""
break break
except json.JSONDecodeError: except json.JSONDecodeError:
logger.bind(tag=TAG).warning("收到无效的JSON消息") logger.bind(tag=TAG).warning("收到无效的JSON消息")
@@ -510,7 +522,6 @@ class TTSProvider(TTSProviderBase):
# 生成会话ID # 生成会话ID
session_id = uuid.uuid4().hex session_id = uuid.uuid4().hex
message_id = uuid.uuid4().hex
# 存储音频数据 # 存储音频数据
audio_data = [] audio_data = []
@@ -529,9 +540,10 @@ class TTSProvider(TTSProviderBase):
) )
try: try:
# 发送StartSynthesis请求 # 发送StartSynthesis请求
start_message_id = str(uuid.uuid4().hex)
start_request = { start_request = {
"header": { "header": {
"message_id": message_id, "message_id": start_message_id,
"task_id": session_id, "task_id": session_id,
"namespace": "FlowingSpeechSynthesizer", "namespace": "FlowingSpeechSynthesizer",
"name": "StartSynthesis", "name": "StartSynthesis",
@@ -544,31 +556,53 @@ class TTSProvider(TTSProviderBase):
"volume": self.volume, "volume": self.volume,
"speech_rate": self.speech_rate, "speech_rate": self.speech_rate,
"pitch_rate": self.pitch_rate, "pitch_rate": self.pitch_rate,
"enable_subtitle": True "enable_subtitle": True,
} },
} }
await ws.send(json.dumps(start_request)) await ws.send(json.dumps(start_request))
# 等待SynthesisStarted响应
synthesis_started = False
while not synthesis_started:
msg = await ws.recv()
if isinstance(msg, str):
data = json.loads(msg)
header = data.get("header", {})
if header.get("name") == "SynthesisStarted":
synthesis_started = True
logger.bind(tag=TAG).debug("TTS合成已启动")
elif header.get("name") == "TaskFailed":
error_info = data.get("payload", {}).get(
"error_info", {}
)
error_code = error_info.get("error_code")
error_message = error_info.get(
"error_message", "未知错误"
)
raise Exception(
f"启动合成失败: {error_code} - {error_message}"
)
# 发送文本合成请求 # 发送文本合成请求
filtered_text = MarkdownCleaner.clean_markdown(text) filtered_text = MarkdownCleaner.clean_markdown(text)
run_message_id = str(uuid.uuid4().hex)
run_request = { run_request = {
"header": { "header": {
"message_id": message_id, "message_id": run_message_id,
"task_id": session_id, "task_id": session_id,
"namespace": "FlowingSpeechSynthesizer", "namespace": "FlowingSpeechSynthesizer",
"name": "RunSynthesis", "name": "RunSynthesis",
"appkey": self.appkey, "appkey": self.appkey,
}, },
"payload": { "payload": {"text": filtered_text},
"text": filtered_text
}
} }
await ws.send(json.dumps(run_request)) await ws.send(json.dumps(run_request))
# 发送停止合成请求 # 发送停止合成请求
stop_message_id = str(uuid.uuid4().hex)
stop_request = { stop_request = {
"header": { "header": {
"message_id": message_id, "message_id": stop_message_id,
"task_id": session_id, "task_id": session_id,
"namespace": "FlowingSpeechSynthesizer", "namespace": "FlowingSpeechSynthesizer",
"name": "StopSynthesis", "name": "StopSynthesis",
@@ -578,17 +612,33 @@ class TTSProvider(TTSProviderBase):
await ws.send(json.dumps(stop_request)) await ws.send(json.dumps(stop_request))
# 接收音频数据 # 接收音频数据
while True: synthesis_completed = False
while not synthesis_completed:
msg = await ws.recv() msg = await ws.recv()
if isinstance(msg, (bytes, bytearray)): if isinstance(msg, (bytes, bytearray)):
# 编码为Opus并收集 # 编码为Opus并收集
opus_frames = self.opus_encoder.encode_pcm_to_opus(msg, False) opus_frames = self.opus_encoder.encode_pcm_to_opus(
msg, False
)
audio_data.extend(opus_frames) audio_data.extend(opus_frames)
elif isinstance(msg, str): elif isinstance(msg, str):
data = json.loads(msg) data = json.loads(msg)
header = data.get("header", {}) header = data.get("header", {})
if header.get("name") == "SynthesisCompleted": event_name = header.get("name")
break if event_name == "SynthesisCompleted":
synthesis_completed = True
logger.bind(tag=TAG).debug("TTS合成完成")
elif event_name == "TaskFailed":
error_info = data.get("payload", {}).get(
"error_info", {}
)
error_code = error_info.get("error_code")
error_message = error_info.get(
"error_message", "未知错误"
)
raise Exception(
f"合成失败: {error_code} - {error_message}"
)
finally: finally:
try: try:
await ws.close() await ws.close()