resolve merge conflict

This commit is contained in:
caixypromise
2025-12-29 00:09:26 +08:00
375 changed files with 45358 additions and 9704 deletions
@@ -0,0 +1,522 @@
import os
import uuid
import json
import time
import queue
import asyncio
import traceback
import websockets
from asyncio import Task
from config.logger import setup_logging
from core.utils import opus_encoder_utils
from core.utils.tts import MarkdownCleaner
from core.providers.tts.base import TTSProviderBase
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
TAG = __name__
logger = setup_logging()
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.interface_type = InterfaceType.DUAL_STREAM
# 基础配置
self.api_key = config.get("api_key")
if not self.api_key:
raise ValueError("api_key is required for CosyVoice TTS")
# WebSocket配置
self.ws_url = "wss://dashscope.aliyuncs.com/api-ws/v1/inference/"
self.ws = None
self._monitor_task = None
self.last_active_time = None
# 模型和音色配置
self.model = config.get("model", "cosyvoice-v2")
self.voice = config.get("voice", "longxiaochun_v2") # 默认音色
if config.get("private_voice"):
self.voice = config.get("private_voice")
# 音频参数配置
self.format = config.get("format", "pcm")
sample_rate = config.get("sample_rate", "24000")
self.sample_rate = int(sample_rate) if sample_rate else 24000
volume = config.get("volume", "50")
self.volume = int(volume) if volume else 50
rate = config.get("rate", "1.0")
self.rate = float(rate) if rate else 1.0
pitch = config.get("pitch", "1.0")
self.pitch = float(pitch) if pitch else 1.0
self.header = {
"Authorization": f"Bearer {self.api_key}",
# "user-agent": "your_platform_info", // 可选
# "X-DashScope-WorkSpace": workspace, // 可选,阿里云百炼业务空间ID
"X-DashScope-DataInspection": "enable",
}
# 创建Opus编码器
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
sample_rate=self.sample_rate, channels=1, frame_size_ms=60
)
async def _ensure_connection(self):
"""确保WebSocket连接可用,支持60秒内连接复用"""
try:
current_time = time.time()
if self.ws and current_time - self.last_active_time < 60:
# 一分钟内才可以复用链接进行连续对话
logger.bind(tag=TAG).info(f"使用已有链接...")
return self.ws
logger.bind(tag=TAG).info("开始建立新连接...")
self.ws = await websockets.connect(
self.ws_url,
additional_headers=self.header,
ping_interval=30,
ping_timeout=10,
close_timeout=10,
)
logger.bind(tag=TAG).info("WebSocket连接建立成功")
self.last_active_time = current_time
return self.ws
except Exception as e:
logger.bind(tag=TAG).error(f"建立连接失败: {str(e)}")
self.ws = None
self.last_active_time = None
raise
def tts_text_priority_thread(self):
"""流式TTS文本处理线程"""
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} | 会话ID: {self.conn.sentence_id}"
)
if message.sentence_type == SentenceType.FIRST:
self.conn.client_abort = False
if self.conn.client_abort:
try:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
continue
except Exception as e:
logger.bind(tag=TAG).error(f"取消TTS会话失败: {str(e)}")
continue
if message.sentence_type == SentenceType.FIRST:
# 初始化会话
try:
if not getattr(self.conn, "sentence_id", None):
self.conn.sentence_id = uuid.uuid4().hex
logger.bind(tag=TAG).info(f"自动生成新的 会话ID: {self.conn.sentence_id}")
logger.bind(tag=TAG).info("开始启动TTS会话...")
future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.sentence_id),
loop=self.conn.loop,
)
future.result()
self.before_stop_play_files.clear()
logger.bind(tag=TAG).info("TTS会话启动成功")
except Exception as e:
logger.bind(tag=TAG).error(f"启动TTS会话失败: {str(e)}")
continue
elif ContentType.TEXT == message.content_type:
if message.content_detail:
try:
logger.bind(tag=TAG).debug(
f"开始发送TTS文本: {message.content_detail}"
)
future = asyncio.run_coroutine_threadsafe(
self.text_to_speak(message.content_detail, None),
loop=self.conn.loop,
)
future.result()
logger.bind(tag=TAG).debug("TTS文本发送成功")
except Exception as e:
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
continue
elif ContentType.FILE == message.content_type:
logger.bind(tag=TAG).info(
f"添加音频文件到待播放列表: {message.content_file}"
)
if message.content_file and os.path.exists(message.content_file):
# 先处理文件音频数据
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
if message.sentence_type == SentenceType.LAST:
try:
logger.bind(tag=TAG).info("开始结束TTS会话...")
future = asyncio.run_coroutine_threadsafe(
self.finish_session(self.conn.sentence_id),
loop=self.conn.loop,
)
future.result()
except Exception as e:
logger.bind(tag=TAG).error(f"结束TTS会话失败: {str(e)}")
continue
except queue.Empty:
continue
except Exception as e:
logger.bind(tag=TAG).error(
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
)
continue
async def text_to_speak(self, text, _):
"""发送文本到TTS服务进行合成"""
try:
if self.ws is None:
logger.bind(tag=TAG).warning("WebSocket连接不存在,终止发送文本")
return
# 过滤Markdown
filtered_text = MarkdownCleaner.clean_markdown(text)
# 发送continue-task消息
continue_task_message = {
"header": {
"action": "continue-task",
"task_id": self.conn.sentence_id,
"streaming": "duplex",
},
"payload": {"input": {"text": filtered_text}},
}
await self.ws.send(json.dumps(continue_task_message))
self.last_active_time = time.time()
logger.bind(tag=TAG).debug(f"已发送文本: {filtered_text}")
except Exception as e:
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
if self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
raise
async def start_session(self, session_id):
"""启动TTS会话"""
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
try:
# 检查并清理上一个会话的监听任务
if (
self._monitor_task is not None
and isinstance(self._monitor_task, Task)
and not self._monitor_task.done()
):
logger.bind(tag=TAG).info("检测到未完成的上个会话,关闭监听任务...")
await self.close()
# 确保连接可用
await self._ensure_connection()
# 启动监听任务
self._monitor_task = asyncio.create_task(self._start_monitor_tts_response())
# 发送run-task消息启动会话
run_task_message = {
"header": {
"action": "run-task",
"task_id": session_id,
"streaming": "duplex",
},
"payload": {
"task_group": "audio",
"task": "tts",
"function": "SpeechSynthesizer",
"model": self.model,
"parameters": {
"text_type": "PlainText",
"voice": self.voice,
"format": self.format,
"sample_rate": self.sample_rate,
"volume": self.volume,
"rate": self.rate,
"pitch": self.pitch,
},
"input": {}
},
}
await self.ws.send(json.dumps(run_task_message))
self.last_active_time = time.time()
logger.bind(tag=TAG).info("会话启动请求已发送")
except Exception as e:
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
await self.close()
raise
async def finish_session(self, session_id):
"""结束TTS会话"""
logger.bind(tag=TAG).info(f"关闭会话~~{session_id}")
try:
if self.ws and session_id:
# 发送finish-task消息
finish_task_message = {
"header": {
"action": "finish-task",
"task_id": session_id,
"streaming": "duplex",
},
"payload": {
"input": {}
}
}
await self.ws.send(json.dumps(finish_task_message))
self.last_active_time = time.time()
logger.bind(tag=TAG).info("会话结束请求已发送")
# 等待监听任务完成
if self._monitor_task:
try:
await self._monitor_task
except Exception as e:
logger.bind(tag=TAG).error(
f"等待监听任务完成时发生错误: {str(e)}"
)
finally:
self._monitor_task = None
except Exception as e:
logger.bind(tag=TAG).error(f"关闭会话失败: {str(e)}")
await self.close()
raise
async def close(self):
"""清理资源"""
# 取消监听任务
if self._monitor_task:
try:
self._monitor_task.cancel()
await self._monitor_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"关闭时取消监听任务错误: {e}")
self._monitor_task = None
# 关闭WebSocket连接
if self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
self.last_active_time = None
async def _start_monitor_tts_response(self):
"""监听TTS响应"""
try:
session_finished = False
while not self.conn.stop_event.is_set():
try:
msg = await self.ws.recv()
self.last_active_time = time.time()
# 检查客户端是否中止
if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止监听TTS响应")
break
if isinstance(msg, str): # JSON控制消息
try:
data = json.loads(msg)
event = data["header"].get("event")
if event == "task-started":
logger.bind(tag=TAG).debug("TTS任务启动成功~")
self.tts_audio_queue.put((SentenceType.FIRST, [], None))
elif event == "result-generated":
# 发送缓存的数据
if self.conn.tts_MessageText:
logger.bind(tag=TAG).info(
f"句子语音生成成功: {self.conn.tts_MessageText}"
)
self.tts_audio_queue.put(
(SentenceType.FIRST, [], self.conn.tts_MessageText)
)
self.conn.tts_MessageText = None
elif event == "task-finished":
logger.bind(tag=TAG).debug("TTS任务完成~")
self._process_before_stop_play_files()
session_finished = True
break
elif event == "task-failed":
error_code = data["header"].get("error_code", "unknown")
error_message = data["header"].get("error_message", "未知错误")
logger.bind(tag=TAG).error(
f"TTS任务失败: {error_code} - {error_message}"
)
break
except json.JSONDecodeError:
logger.bind(tag=TAG).warning("收到无效的JSON消息")
elif isinstance(msg, (bytes, bytearray)):
self.opus_encoder.encode_pcm_to_opus_stream(
msg, False, callback=self.handle_opus
)
except websockets.ConnectionClosed:
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
break
except Exception as e:
logger.bind(tag=TAG).error(
f"处理TTS响应时出错: {e}\n{traceback.format_exc()}"
)
break
# 仅在连接异常且非正常结束时才关闭连接
if not session_finished and self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
# 监听任务退出时清理引用
finally:
self._monitor_task = None
def to_tts(self, text: str) -> list:
"""非流式生成音频数据,用于生成音频及测试场景"""
try:
# 创建事件循环
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
# 生成会话ID
session_id = uuid.uuid4().hex
# 存储音频数据
audio_data = []
async def _generate_audio():
ws = await websockets.connect(
self.ws_url,
additional_headers=self.header,
ping_interval=30,
ping_timeout=10,
close_timeout=10,
max_size=10 * 1024 * 1024,
)
try:
# 发送run-task消息启动会话
run_task_message = {
"header": {
"action": "run-task",
"task_id": session_id,
"streaming": "duplex",
},
"payload": {
"task_group": "audio",
"task": "tts",
"function": "SpeechSynthesizer",
"model": self.model,
"parameters": {
"text_type": "PlainText",
"voice": self.voice,
"format": self.format,
"sample_rate": self.sample_rate,
"volume": self.volume,
"rate": self.rate,
"pitch": self.pitch,
},
"input": {}
},
}
await ws.send(json.dumps(run_task_message))
# 等待任务启动
task_started = False
while not task_started:
msg = await ws.recv()
if isinstance(msg, str):
data = json.loads(msg)
header = data.get("header", {})
if header.get("event") == "task-started":
task_started = True
logger.bind(tag=TAG).debug("TTS任务已启动")
elif header.get("event") == "task-failed":
error_code = header.get("error_code", "unknown")
error_message = header.get("error_message", "未知错误")
raise Exception(
f"启动任务失败: {error_code} - {error_message}"
)
# 发送文本
filtered_text = MarkdownCleaner.clean_markdown(text)
# 发送continue-task消息
continue_task_message = {
"header": {
"action": "continue-task",
"task_id": session_id,
"streaming": "duplex",
},
"payload": {"input": {"text": filtered_text}},
}
await ws.send(json.dumps(continue_task_message))
# 发送finish-task消息
finish_task_message = {
"header": {
"action": "finish-task",
"task_id": session_id,
"streaming": "duplex",
},
"payload": {
"input": {}
}
}
await ws.send(json.dumps(finish_task_message))
# 接收音频数据
task_finished = False
while not task_finished:
msg = await ws.recv()
if isinstance(msg, (bytes, bytearray)):
self.opus_encoder.encode_pcm_to_opus_stream(
msg,
end_of_stream=False,
callback=lambda opus: audio_data.append(opus)
)
elif isinstance(msg, str):
data = json.loads(msg)
header = data.get("header", {})
if header.get("event") == "task-finished":
task_finished = True
logger.bind(tag=TAG).debug("TTS任务完成")
elif header.get("event") == "task-failed":
error_code = header.get("error_code", "unknown")
error_message = header.get("error_message", "未知错误")
raise Exception(
f"合成失败: {error_code} - {error_message}"
)
finally:
# 清理资源
try:
await ws.close()
except:
pass
# 运行异步任务
loop.run_until_complete(_generate_audio())
loop.close()
return audio_data
except Exception as e:
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
return []
@@ -1,3 +1,4 @@
import random
import uuid
import json
import hmac
@@ -131,7 +132,7 @@ class TTSProvider(TTSProviderBase):
self.last_active_time = None
# 专属tts设置
self.message_id = ""
self.task_id = uuid.uuid4().hex
# 创建Opus编码器
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
@@ -185,9 +186,10 @@ class TTSProvider(TTSProviderBase):
current_time = time.time()
if self.ws and current_time - self.last_active_time < 10:
# 10秒内才可以复用链接进行连续对话
logger.bind(tag=TAG).info(f"使用已有链接...")
self.task_id = uuid.uuid4().hex
logger.bind(tag=TAG).info(f"使用已有链接..., task_id: {self.task_id}")
return self.ws
logger.bind(tag=TAG).info("开始建立新连接...")
logger.bind(tag=TAG).debug("开始建立新连接...")
self.ws = await websockets.connect(
self.ws_url,
@@ -196,7 +198,8 @@ class TTSProvider(TTSProviderBase):
ping_timeout=10,
close_timeout=10,
)
logger.bind(tag=TAG).info("WebSocket连接建立成功")
self.task_id = uuid.uuid4().hex
logger.bind(tag=TAG).debug(f"WebSocket连接建立成功, task_id: {self.task_id}")
self.last_active_time = time.time()
return self.ws
except Exception as e:
@@ -224,23 +227,14 @@ class TTSProvider(TTSProviderBase):
if message.sentence_type == SentenceType.FIRST:
# 初始化参数
try:
if not getattr(self.conn, "sentence_id", None):
self.conn.sentence_id = uuid.uuid4().hex
logger.bind(tag=TAG).info(
f"自动生成新的 会话ID: {self.conn.sentence_id}"
)
# aliyunStream独有的参数生成
self.message_id = str(uuid.uuid4().hex)
logger.bind(tag=TAG).info("开始启动TTS会话...")
logger.bind(tag=TAG).debug("开始启动TTS会话...")
future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.sentence_id),
self.start_session(self.task_id),
loop=self.conn.loop,
)
future.result()
self.before_stop_play_files.clear()
logger.bind(tag=TAG).info("TTS会话启动成功")
logger.bind(tag=TAG).debug("TTS会话启动成功")
except Exception as e:
logger.bind(tag=TAG).error(f"启动TTS会话失败: {str(e)}")
@@ -271,9 +265,9 @@ class TTSProvider(TTSProviderBase):
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
if message.sentence_type == SentenceType.LAST:
try:
logger.bind(tag=TAG).info("开始结束TTS会话...")
logger.bind(tag=TAG).debug("开始结束TTS会话...")
future = asyncio.run_coroutine_threadsafe(
self.finish_session(self.conn.sentence_id),
self.finish_session(self.task_id),
loop=self.conn.loop,
)
future.result()
@@ -296,8 +290,8 @@ class TTSProvider(TTSProviderBase):
filtered_text = MarkdownCleaner.clean_markdown(text)
run_request = {
"header": {
"message_id": self.message_id,
"task_id": self.conn.sentence_id,
"message_id": uuid.uuid4().hex,
"task_id": self.task_id,
"namespace": "FlowingSpeechSynthesizer",
"name": "RunSynthesis",
"appkey": self.appkey,
@@ -318,8 +312,8 @@ class TTSProvider(TTSProviderBase):
self.ws = None
raise
async def start_session(self, session_id):
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
async def start_session(self, task_id):
logger.bind(tag=TAG).debug("开始会话~~")
try:
# 会话开始时检测上个会话的监听状态
if (
@@ -340,8 +334,8 @@ class TTSProvider(TTSProviderBase):
start_request = {
"header": {
"message_id": self.message_id,
"task_id": self.conn.sentence_id,
"message_id": uuid.uuid4().hex,
"task_id": self.task_id,
"namespace": "FlowingSpeechSynthesizer",
"name": "StartSynthesis",
"appkey": self.appkey,
@@ -358,28 +352,28 @@ class TTSProvider(TTSProviderBase):
}
await self.ws.send(json.dumps(start_request))
self.last_active_time = time.time()
logger.bind(tag=TAG).info("会话启动请求已发送")
logger.bind(tag=TAG).debug("会话启动请求已发送")
except Exception as e:
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
# 确保清理资源
await self.close()
raise
async def finish_session(self, session_id):
logger.bind(tag=TAG).info(f"关闭会话~~{session_id}")
async def finish_session(self, task_id):
logger.bind(tag=TAG).debug(f"关闭会话~~{task_id}")
try:
if self.ws:
stop_request = {
"header": {
"message_id": self.message_id,
"task_id": self.conn.sentence_id,
"message_id": uuid.uuid4().hex,
"task_id": self.task_id,
"namespace": "FlowingSpeechSynthesizer",
"name": "StopSynthesis",
"appkey": self.appkey,
}
}
await self.ws.send(json.dumps(stop_request))
logger.bind(tag=TAG).info("会话结束请求已发送")
logger.bind(tag=TAG).debug("会话结束请求已发送")
self.last_active_time = time.time()
if self._monitor_task:
try:
@@ -457,7 +451,6 @@ class TTSProvider(TTSProviderBase):
logger.bind(tag=TAG).warning("收到无效的JSON消息")
# 二进制消息(音频数据)
elif isinstance(msg, (bytes, bytearray)):
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
self.opus_encoder.encode_pcm_to_opus_stream(msg, False, self.handle_opus)
except websockets.ConnectionClosed:
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
@@ -485,8 +478,6 @@ class TTSProvider(TTSProviderBase):
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
# 生成会话ID
session_id = uuid.uuid4().hex
# 存储音频数据
audio_data = []
@@ -505,11 +496,10 @@ class TTSProvider(TTSProviderBase):
)
try:
# 发送StartSynthesis请求
start_message_id = str(uuid.uuid4().hex)
start_request = {
"header": {
"message_id": start_message_id,
"task_id": session_id,
"message_id": uuid.uuid4().hex,
"task_id": self.task_id,
"namespace": "FlowingSpeechSynthesizer",
"name": "StartSynthesis",
"appkey": self.appkey,
@@ -550,11 +540,10 @@ class TTSProvider(TTSProviderBase):
# 发送文本合成请求
filtered_text = MarkdownCleaner.clean_markdown(text)
run_message_id = str(uuid.uuid4().hex)
run_request = {
"header": {
"message_id": run_message_id,
"task_id": session_id,
"message_id": uuid.uuid4().hex,
"task_id": self.task_id,
"namespace": "FlowingSpeechSynthesizer",
"name": "RunSynthesis",
"appkey": self.appkey,
@@ -564,11 +553,10 @@ class TTSProvider(TTSProviderBase):
await ws.send(json.dumps(run_request))
# 发送停止合成请求
stop_message_id = str(uuid.uuid4().hex)
stop_request = {
"header": {
"message_id": stop_message_id,
"task_id": session_id,
"message_id": uuid.uuid4().hex,
"task_id": self.task_id,
"namespace": "FlowingSpeechSynthesizer",
"name": "StopSynthesis",
"appkey": self.appkey,
@@ -395,14 +395,6 @@ class TTSProviderBase(ABC):
# 收到下一个文本开始或会话结束时进行上报
if sentence_type is not SentenceType.MIDDLE:
# 重置音频流控状态(新句子开始或者结束)
if hasattr(self.conn, 'audio_flow_control'):
self.conn.audio_flow_control = {
'last_send_time': 0,
'packet_count': 0,
'start_time': time.perf_counter()
}
# 上报TTS数据
if enqueue_text is not None and enqueue_audio is not None:
enqueue_tts_report(self.conn, enqueue_text, enqueue_audio)
@@ -149,6 +149,7 @@ class TTSProvider(TTSProviderBase):
self.access_token = config.get("access_token")
self.cluster = config.get("cluster")
self.resource_id = config.get("resource_id")
self.activate_session = False
if config.get("private_voice"):
self.voice = config.get("private_voice")
else:
@@ -162,7 +163,8 @@ class TTSProvider(TTSProviderBase):
self.ws_url = config.get("ws_url")
self.authorization = config.get("authorization")
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
self.enable_two_way = True
enable_ws_reuse_value = config.get("enable_ws_reuse", True)
self.enable_ws_reuse = False if str(enable_ws_reuse_value).lower() in ('false', 'False') else True
self.tts_text = ""
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
sample_rate=16000, channels=1, frame_size_ms=60
@@ -180,12 +182,18 @@ class TTSProvider(TTSProviderBase):
raise
async def _ensure_connection(self):
"""建立新的WebSocket连接"""
"""建立新的WebSocket连接,并启动监听任务(仅第一次)"""
try:
if self.ws:
logger.bind(tag=TAG).info(f"使用已有链接...")
return self.ws
logger.bind(tag=TAG).info("开始建立新连接...")
if self.enable_ws_reuse:
logger.bind(tag=TAG).info(f"使用已有链接...")
return self.ws
else:
try:
await self.finish_connection()
except:
pass
logger.bind(tag=TAG).debug("开始建立新连接...")
ws_header = {
"X-Api-App-Key": self.appId,
"X-Api-Access-Key": self.access_token,
@@ -195,12 +203,34 @@ class TTSProvider(TTSProviderBase):
self.ws = await websockets.connect(
self.ws_url, additional_headers=ws_header, max_size=1000000000
)
logger.bind(tag=TAG).info("WebSocket连接建立成功")
logger.bind(tag=TAG).debug("WebSocket连接建立成功")
# 连接建立成功后,启动监听任务
if self._monitor_task is None or self._monitor_task.done():
logger.bind(tag=TAG).debug("启动监听任务...")
self._monitor_task = asyncio.create_task(self._start_monitor_tts_response())
return self.ws
except Exception as e:
logger.bind(tag=TAG).error(f"建立连接失败: {str(e)}")
self.ws = None
raise
async def finish_connection(self):
"""发送 FinishConnection 事件,等待服务端返回 EVENT_ConnectionFinished"""
try:
if self.ws:
logger.bind(tag=TAG).debug("开始关闭连接...")
header = Header(
message_type=FULL_CLIENT_REQUEST,
message_type_specific_flags=MsgTypeFlagWithEvent,
serial_method=JSON,
).as_bytes()
optional = Optional(event=EVENT_FinishConnection).as_bytes()
payload = str.encode("{}")
await self.send_event(self.ws, header, optional, payload)
except:
pass
def tts_text_priority_thread(self):
"""火山引擎双流式TTS的文本处理线程"""
@@ -217,10 +247,16 @@ class TTSProvider(TTSProviderBase):
if self.conn.client_abort:
try:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
asyncio.run_coroutine_threadsafe(
self.cancel_session(self.conn.sentence_id),
loop=self.conn.loop,
)
if self.enable_ws_reuse:
asyncio.run_coroutine_threadsafe(
self.cancel_session(self.conn.sentence_id),
loop=self.conn.loop,
)
else:
asyncio.run_coroutine_threadsafe(
self.finish_connection(),
loop=self.conn.loop,
)
continue
except Exception as e:
logger.bind(tag=TAG).error(f"取消TTS会话失败: {str(e)}")
@@ -231,16 +267,16 @@ class TTSProvider(TTSProviderBase):
try:
if not getattr(self.conn, "sentence_id", None):
self.conn.sentence_id = uuid.uuid4().hex
logger.bind(tag=TAG).info(f"自动生成新的 会话ID: {self.conn.sentence_id}")
logger.bind(tag=TAG).debug(f"自动生成新的 会话ID: {self.conn.sentence_id}")
logger.bind(tag=TAG).info("开始启动TTS会话...")
logger.bind(tag=TAG).debug("开始启动TTS会话...")
future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.sentence_id),
loop=self.conn.loop,
)
future.result()
self.before_stop_play_files.clear()
logger.bind(tag=TAG).info("TTS会话启动成功")
logger.bind(tag=TAG).debug("TTS会话启动成功")
except Exception as e:
logger.bind(tag=TAG).error(f"启动TTS会话失败: {str(e)}")
continue
@@ -270,7 +306,7 @@ class TTSProvider(TTSProviderBase):
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
if message.sentence_type == SentenceType.LAST:
try:
logger.bind(tag=TAG).info("开始结束TTS会话...")
logger.bind(tag=TAG).debug("开始结束TTS会话...")
future = asyncio.run_coroutine_threadsafe(
self.finish_session(self.conn.sentence_id),
loop=self.conn.loop,
@@ -313,23 +349,25 @@ class TTSProvider(TTSProviderBase):
raise
async def start_session(self, session_id):
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
try:
# 会话开始时检测上个会话的监听状态
if (
self._monitor_task is not None
and isinstance(self._monitor_task, Task)
and not self._monitor_task.done()
):
logger.bind(tag=TAG).info("检测到未完成的上个会话,关闭监听任务和连接...")
logger.bind(tag=TAG).debug(f"开始会话~~{session_id}")
try:
# 等待上一个会话结束,最多等待3次
for _ in range(3):
if not self.activate_session:
break
logger.bind(tag=TAG).debug(f"等待上一个会话结束...")
await asyncio.sleep(0.1)
else:
# 等待超时,强制清除连接状态
logger.bind(tag=TAG).debug("等待上一个会话超时,清除连接状态...")
await self.close()
# 建立新连接
# 设置会话激活标志
self.activate_session = True
# 确保连接建立
await self._ensure_connection()
# 启动监听任务
self._monitor_task = asyncio.create_task(self._start_monitor_tts_response())
header = Header(
message_type=FULL_CLIENT_REQUEST,
message_type_specific_flags=MsgTypeFlagWithEvent,
@@ -342,7 +380,7 @@ class TTSProvider(TTSProviderBase):
event=EVENT_StartSession, speaker=self.voice
)
await self.send_event(self.ws, header, optional, payload)
logger.bind(tag=TAG).info("会话启动请求已发送")
logger.bind(tag=TAG).debug("会话启动请求已发送")
except Exception as e:
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
# 确保清理资源
@@ -350,7 +388,7 @@ class TTSProvider(TTSProviderBase):
raise
async def finish_session(self, session_id):
logger.bind(tag=TAG).info(f"关闭会话~~{session_id}")
logger.bind(tag=TAG).debug(f"关闭会话~~{session_id}")
try:
if self.ws:
header = Header(
@@ -363,18 +401,7 @@ class TTSProvider(TTSProviderBase):
).as_bytes()
payload = str.encode("{}")
await self.send_event(self.ws, header, optional, payload)
logger.bind(tag=TAG).info("会话结束请求已发送")
# 等待监听任务完成
if self._monitor_task:
try:
await self._monitor_task
except Exception as e:
logger.bind(tag=TAG).error(
f"等待监听任务完成时发生错误: {str(e)}"
)
finally:
self._monitor_task = None
logger.bind(tag=TAG).debug("会话结束请求已发送")
except Exception as e:
logger.bind(tag=TAG).error(f"关闭会话失败: {str(e)}")
@@ -383,7 +410,7 @@ class TTSProvider(TTSProviderBase):
raise
async def cancel_session(self,session_id):
logger.bind(tag=TAG).info(f"取消会话,释放服务端资源~~{session_id}")
logger.bind(tag=TAG).debug(f"取消会话,释放服务端资源~~{session_id}")
try:
if self.ws:
header = Header(
@@ -396,7 +423,7 @@ class TTSProvider(TTSProviderBase):
).as_bytes()
payload = str.encode("{}")
await self.send_event(self.ws, header, optional, payload)
logger.bind(tag=TAG).info("会话取消请求已发送")
logger.bind(tag=TAG).debug("会话取消请求已发送")
except Exception as e:
logger.bind(tag=TAG).error(f"取消会话失败: {str(e)}")
# 确保清理资源
@@ -405,6 +432,7 @@ class TTSProvider(TTSProviderBase):
async def close(self):
"""资源清理方法"""
self.activate_session = False
# 取消监听任务
if self._monitor_task:
try:
@@ -424,9 +452,8 @@ class TTSProvider(TTSProviderBase):
self.ws = None
async def _start_monitor_tts_response(self):
"""监听TTS响应"""
"""监听TTS响应 - 长期运行"""
try:
session_finished = False # 标记会话是否正常结束
while not self.conn.stop_event.is_set():
try:
# 确保 `recv()` 运行在同一个 event loop
@@ -434,10 +461,22 @@ class TTSProvider(TTSProviderBase):
res = self.parser_response(msg)
self.print_response(res, "send_text res:")
# 优先处理连接级别事件
if res.optional.event == EVENT_ConnectionFinished:
logger.bind(tag=TAG).debug(f"链接关闭成功~~")
break
# 只处理当前活跃会话的响应
if res.optional.sessionId and self.conn.sentence_id != res.optional.sessionId:
# 如果是会话结束相关事件,即使会话ID不匹配也要重置状态
if res.optional.event in [EVENT_SessionCanceled, EVENT_SessionFailed, EVENT_SessionFinished]:
logger.bind(tag=TAG).debug(f"收到残余下行结束响应重置会话状态~~")
self.activate_session = False
continue
if res.optional.event == EVENT_SessionCanceled:
logger.bind(tag=TAG).debug(f"释放服务端资源成功~~")
session_finished = True
break
self.activate_session = False
elif res.optional.event == EVENT_TTSSentenceStart:
json_data = json.loads(res.payload.decode("utf-8"))
self.tts_text = json_data.get("text", "")
@@ -449,15 +488,16 @@ class TTSProvider(TTSProviderBase):
res.optional.event == EVENT_TTSResponse
and res.header.message_type == AUDIO_ONLY_RESPONSE
):
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
self.wav_to_opus_data_audio_raw_stream(res.payload, callback=self.handle_opus)
elif res.optional.event == EVENT_TTSSentenceEnd:
logger.bind(tag=TAG).info(f"句子语音生成成功:{self.tts_text}")
elif res.optional.event == EVENT_SessionFinished:
logger.bind(tag=TAG).debug(f"会话结束~~")
self.activate_session = False
self._process_before_stop_play_files()
session_finished = True
break
# 非复用模式下,会话结束后发送 FinishConnection
if not self.enable_ws_reuse:
await self.finish_connection()
except websockets.ConnectionClosed:
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
break
@@ -467,8 +507,8 @@ class TTSProvider(TTSProviderBase):
)
traceback.print_exc()
break
# 仅在连接异常时关闭
if not session_finished and self.ws:
# 连接异常时关闭WebSocket
if self.ws:
try:
await self.ws.close()
except:
@@ -476,6 +516,7 @@ class TTSProvider(TTSProviderBase):
self.ws = None
# 监听任务退出时清理引用
finally:
self.activate_session = False
self._monitor_task = None
async def send_event(
@@ -514,7 +555,7 @@ class TTSProvider(TTSProviderBase):
def read_res_content(self, res: bytes, offset: int):
content_size = int.from_bytes(res[offset : offset + 4], "big", signed=True)
offset += 4
content = str(res[offset : offset + content_size])
content = res[offset : offset + content_size].decode('utf-8')
offset += content_size
return content, offset
@@ -616,12 +657,13 @@ class TTSProvider(TTSProviderBase):
"speech_rate": self.speech_rate,
"loudness_rate": self.loudness_rate
},
"additions": json.dumps({
"post_process": {
"pitch": self.pitch
}
})
},
"additions": {
"post_process": {
"pitch": self.pitch
}
}
}
)
)
@@ -190,7 +190,7 @@ class TTSProvider(TTSProviderBase):
start_time = time.time()
text = MarkdownCleaner.clean_markdown(text)
payload = {"text": text, "character": self.character}
payload = {"text": text, "character": self.voice}
try:
with requests.post(self.api_url, json=payload, timeout=5) as response:
@@ -1,95 +0,0 @@
import os
import uuid
import json
import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
from core.utils.util import parse_string_to_list
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.group_id = config.get("group_id")
self.api_key = config.get("api_key")
self.model = config.get("model")
if config.get("private_voice"):
self.voice = config.get("private_voice")
else:
self.voice = config.get("voice_id")
default_voice_setting = {
"voice_id": "female-shaonv",
"speed": 1,
"vol": 1,
"pitch": 0,
"emotion": "happy",
}
default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]}
defult_audio_setting = {
"sample_rate": 32000,
"bitrate": 128000,
"format": "mp3",
"channel": 1,
}
self.voice_setting = {
**default_voice_setting,
**config.get("voice_setting", {}),
}
self.pronunciation_dict = {
**default_pronunciation_dict,
**config.get("pronunciation_dict", {}),
}
self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})}
self.timber_weights = parse_string_to_list(config.get("timber_weights"))
if self.voice:
self.voice_setting["voice_id"] = self.voice
self.host = "api.minimax.chat"
self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}"
self.header = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}",
}
self.audio_file_type = defult_audio_setting.get("format", "mp3")
def generate_filename(self, extension=".mp3"):
return os.path.join(
self.output_file,
f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
)
async def text_to_speak(self, text, output_file):
request_json = {
"model": self.model,
"text": text,
"stream": False,
"voice_setting": self.voice_setting,
"pronunciation_dict": self.pronunciation_dict,
"audio_setting": self.audio_setting,
}
if type(self.timber_weights) is list and len(self.timber_weights) > 0:
request_json["timber_weights"] = self.timber_weights
request_json["voice_setting"]["voice_id"] = ""
try:
resp = requests.post(
self.api_url, json.dumps(request_json), headers=self.header
)
# 检查返回请求数据的status_code是否为0
if resp.json()["base_resp"]["status_code"] == 0:
data = resp.json()["data"]["audio"]
audio_bytes = bytes.fromhex(data)
if output_file:
with open(output_file, "wb") as file_to_save:
file_to_save.write(audio_bytes)
else:
return audio_bytes
else:
raise Exception(
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
)
except Exception as e:
raise Exception(f"{__name__} error: {e}")
@@ -1,11 +1,20 @@
import os
import uuid
import json
import time
import queue
import asyncio
import aiohttp
import requests
from datetime import datetime
from typing import Iterator, Optional, Union
from core.providers.tts.base import TTSProviderBase
import traceback
from config.logger import setup_logging
from core.utils.tts import MarkdownCleaner
from core.utils.util import parse_string_to_list
from core.providers.tts.base import TTSProviderBase
from core.utils import opus_encoder_utils, textUtils
from core.providers.tts.dto.dto import SentenceType, ContentType
TAG = __name__
logger = setup_logging()
class TTSProvider(TTSProviderBase):
@@ -28,9 +37,9 @@ class TTSProvider(TTSProviderBase):
}
default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]}
defult_audio_setting = {
"sample_rate": 32000,
"sample_rate": 24000,
"bitrate": 128000,
"format": "mp3",
"format": "pcm",
"channel": 1,
}
self.voice_setting = {
@@ -47,66 +56,101 @@ class TTSProvider(TTSProviderBase):
if self.voice:
self.voice_setting["voice_id"] = self.voice
self.host = "api.minimax.chat"
self.host = "api.minimaxi.com" # 备用地址:api-bj.minimaxi.com
self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}"
self.header = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}",
}
self.audio_file_type = defult_audio_setting.get("format", "mp3")
self.audio_file_type = defult_audio_setting.get("format", "pcm")
def generate_filename(self, extension=".mp3"):
return os.path.join(
self.output_file,
f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
sample_rate=24000, channels=1, frame_size_ms=60
)
async def text_to_speak(self, text, output_file):
"""非流式语音合成(保留原有实现)"""
request_json = {
"model": self.model,
"text": text,
"stream": False,
"voice_setting": self.voice_setting,
"pronunciation_dict": self.pronunciation_dict,
"audio_setting": self.audio_setting,
}
# PCM缓冲区
self.pcm_buffer = bytearray()
if type(self.timber_weights) is list and len(self.timber_weights) > 0:
request_json["timber_weights"] = self.timber_weights
request_json["voice_setting"]["voice_id"] = ""
def tts_text_priority_thread(self):
"""流式文本处理线程"""
while not self.conn.stop_event.is_set():
try:
message = self.tts_text_queue.get(timeout=1)
if message.sentence_type == SentenceType.FIRST:
# 初始化参数
self.tts_stop_request = False
self.processed_chars = 0
self.tts_text_buff = []
self.before_stop_play_files.clear()
elif ContentType.TEXT == message.content_type:
self.tts_text_buff.append(message.content_detail)
segment_text = self._get_segment_text()
if segment_text:
self.to_tts_single_stream(segment_text)
try:
resp = requests.post(
self.api_url, json.dumps(request_json), headers=self.header
)
if resp.json()["base_resp"]["status_code"] == 0:
data = resp.json()["data"]["audio"]
audio_bytes = bytes.fromhex(data)
if output_file:
with open(output_file, "wb") as file_to_save:
file_to_save.write(audio_bytes)
else:
return audio_bytes
elif ContentType.FILE == message.content_type:
logger.bind(tag=TAG).info(
f"添加音频文件到待播放列表: {message.content_file}"
)
if message.content_file and os.path.exists(message.content_file):
# 先处理文件音频数据
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
if message.sentence_type == SentenceType.LAST:
# 处理剩余的文本
self._process_remaining_text_stream(True)
except queue.Empty:
continue
except Exception as e:
logger.bind(tag=TAG).error(
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
)
def _process_remaining_text_stream(self, is_last=False):
"""处理剩余的文本并生成语音
Returns:
bool: 是否成功处理了文本
"""
full_text = "".join(self.tts_text_buff)
remaining_text = full_text[self.processed_chars :]
if remaining_text:
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
if segment_text:
self.to_tts_single_stream(segment_text, is_last)
self.processed_chars += len(full_text)
else:
raise Exception(
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
self._process_before_stop_play_files()
else:
self._process_before_stop_play_files()
def to_tts_single_stream(self, text, is_last=False):
try:
max_repeat_time = 5
text = MarkdownCleaner.clean_markdown(text)
try:
asyncio.run(self.text_to_speak(text, is_last))
except Exception as e:
logger.bind(tag=TAG).warning(
f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
)
max_repeat_time -= 1
if max_repeat_time > 0:
logger.bind(tag=TAG).info(
f"语音生成成功: {text},重试{5 - max_repeat_time}"
)
else:
logger.bind(tag=TAG).error(
f"语音生成失败: {text},请检查网络或服务是否正常"
)
except Exception as e:
raise Exception(f"{__name__} error: {e}")
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
finally:
return None
def text_to_speak_stream(
self,
text: str,
chunk_callback: Optional[callable] = None
) -> Iterator[bytes]:
"""
流式语音合成方法
:param text: 要合成的文本
:param chunk_callback: 可选的回调函数,用于处理每个音频块
:return: 生成器,每次产生一个音频数据块(bytes)
"""
request_json = {
async def text_to_speak(self, text, is_last):
"""流式处理TTS音频,每句只推送一次音频列表"""
payload = {
"model": self.model,
"text": text,
"stream": True,
@@ -115,116 +159,183 @@ class TTSProvider(TTSProviderBase):
"audio_setting": self.audio_setting,
}
if isinstance(self.timber_weights, list) and len(self.timber_weights) > 0:
request_json["timber_weights"] = self.timber_weights
request_json["voice_setting"]["voice_id"] = ""
if type(self.timber_weights) is list and len(self.timber_weights) > 0:
payload["timber_weights"] = self.timber_weights
payload["voice_setting"]["voice_id"] = ""
frame_bytes = int(
self.opus_encoder.sample_rate
* self.opus_encoder.channels # 1
* self.opus_encoder.frame_size_ms
/ 1000
* 2
) # 16-bit = 2 bytes
try:
async with aiohttp.ClientSession() as session:
async with session.post(
self.api_url,
headers=self.header,
data=json.dumps(payload),
timeout=10,
) as resp:
if resp.status != 200:
logger.bind(tag=TAG).error(
f"TTS请求失败: {resp.status}, {await resp.text()}"
)
self.tts_audio_queue.put((SentenceType.LAST, [], None))
return
self.pcm_buffer.clear()
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
# 处理音频流数据
buffer = b""
async for chunk in resp.content.iter_any():
if not chunk:
continue
buffer += chunk
while True:
# 查找数据块分隔符
header_pos = buffer.find(b"data: ")
if header_pos == -1:
break
end_pos = buffer.find(b"\n\n", header_pos)
if end_pos == -1:
break
# 提取单个完整JSON块
json_str = buffer[header_pos + 6 : end_pos].decode("utf-8")
buffer = buffer[end_pos + 2 :]
try:
data = json.loads(json_str)
status = data.get("data", {}).get("status", 1)
audio_hex = data.get("data", {}).get("audio")
# 仅处理status=1的有效音频块 忽略status=2的结束汇总块
if status == 1 and audio_hex:
pcm_data = bytes.fromhex(audio_hex)
self.pcm_buffer.extend(pcm_data)
except json.JSONDecodeError as e:
logger.bind(tag=TAG).error(f"JSON解析失败: {e}")
continue
while len(self.pcm_buffer) >= frame_bytes:
frame = bytes(self.pcm_buffer[:frame_bytes])
del self.pcm_buffer[:frame_bytes]
self.opus_encoder.encode_pcm_to_opus_stream(
frame, end_of_stream=False, callback=self.handle_opus
)
# flush 剩余不足一帧的数据
if self.pcm_buffer:
self.opus_encoder.encode_pcm_to_opus_stream(
bytes(self.pcm_buffer),
end_of_stream=True,
callback=self.handle_opus,
)
self.pcm_buffer.clear()
# 如果是最后一段,输出音频获取完毕
if is_last:
self._process_before_stop_play_files()
except Exception as e:
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
self.tts_audio_queue.put((SentenceType.LAST, [], None))
async def close(self):
"""资源清理"""
await super().close()
if hasattr(self, "opus_encoder"):
self.opus_encoder.close()
def to_tts(self, text: str) -> list:
"""非流式TTS处理,用于测试及保存音频文件的场景
Args:
text: 要转换的文本
Returns:
list: 返回opus编码后的音频数据列表
"""
start_time = time.time()
text = MarkdownCleaner.clean_markdown(text)
payload = {
"model": self.model,
"text": text,
"stream": True,
"voice_setting": self.voice_setting,
"pronunciation_dict": self.pronunciation_dict,
"audio_setting": self.audio_setting,
}
if type(self.timber_weights) is list and len(self.timber_weights) > 0:
payload["timber_weights"] = self.timber_weights
payload["voice_setting"]["voice_id"] = ""
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}",
}
try:
with requests.post(
self.api_url,
data=json.dumps(request_json),
headers=self.header,
stream=True
self.api_url, data=json.dumps(payload), headers=headers, timeout=5
) as response:
# 检查HTTP状态码
if response.status_code != 200:
raise Exception(
f"HTTP error: {response.status_code}, response: {response.text}"
logger.bind(tag=TAG).error(
f"TTS请求失败: {response.status_code}, {response.text}"
)
# 处理流式响应
for line in response.iter_lines():
if line: # 过滤空行
# 检查是否为数据行 (SSE格式)
if line.startswith(b'data:'):
try:
data = json.loads(line[5:].strip()) # 去掉"data:"前缀
# 检查API状态码
if data.get("base_resp", {}).get("status_code", -1) != 0:
raise Exception(
f"API error: {data.get('base_resp', {}).get('status_msg')}"
)
# 跳过非音频数据块
if "extra_info" in data:
continue
# 提取音频数据
audio_hex = data.get("data", {}).get("audio")
if audio_hex:
audio_chunk = bytes.fromhex(audio_hex)
if chunk_callback:
chunk_callback(audio_chunk)
yield audio_chunk
except json.JSONDecodeError:
# 忽略JSON解析错误(可能是心跳包等)
continue
except Exception as e:
raise e
except Exception as e:
raise Exception(f"{__name__} stream error: {e}")
return []
def save_stream_to_file(
self,
text: str,
output_file: Optional[str] = None,
progress_callback: Optional[callable] = None
) -> str:
"""
流式合成并保存到文件
:param text: 要合成的文本
:param output_file: 输出文件路径,如果为None则自动生成
:param progress_callback: 可选的回调函数,接收已写入的字节数
:return: 保存的文件路径
"""
if not output_file:
output_file = self.generate_filename(extension=f".{self.audio_file_type}")
os.makedirs(os.path.dirname(output_file), exist_ok=True)
total_bytes = 0
try:
with open(output_file, "wb") as audio_file:
for audio_chunk in self.text_to_speak_stream(text):
audio_file.write(audio_chunk)
audio_file.flush()
total_bytes += len(audio_chunk)
if progress_callback:
progress_callback(total_bytes)
return output_file
except Exception as e:
# 清理可能创建的不完整文件
if os.path.exists(output_file):
os.remove(output_file)
raise e
logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}")
# 使用opus编码器处理PCM数据
opus_datas = []
full_content = response.content.decode('utf-8')
pcm_data = bytearray()
for data_block in full_content.split('\n\n'):
if not data_block.startswith('data: '):
continue
try:
json_str = data_block[6:] # 去除'data: '前缀
data = json.loads(json_str)
if data.get('data', {}).get('status') == 1:
audio_hex = data['data']['audio']
pcm_data.extend(bytes.fromhex(audio_hex))
except (json.JSONDecodeError, KeyError) as e:
logger.bind(tag=TAG).warning(f"无效数据块: {e}")
continue
# 计算每帧的字节数
frame_bytes = int(
self.opus_encoder.sample_rate
* self.opus_encoder.channels
* self.opus_encoder.frame_size_ms
/ 1000
* 2
)
# 分帧处理合并后的PCM数据
for i in range(0, len(pcm_data), frame_bytes):
frame = bytes(pcm_data[i:i+frame_bytes])
if len(frame) < frame_bytes:
frame += b"\x00" * (frame_bytes - len(frame))
self.opus_encoder.encode_pcm_to_opus_stream(
frame,
end_of_stream=(i + frame_bytes >= len(pcm_data)),
callback=lambda opus: opus_datas.append(opus)
)
return opus_datas
def stream_to_audio_player(self, text: str, player_command: list = None):
"""
流式合成并直接播放音频
:param text: 要合成的文本
:param player_command: 音频播放器命令,默认使用mpv
"""
if player_command is None:
player_command = ["mpv", "--no-cache", "--no-terminal", "--", "fd://0"]
try:
import subprocess
player_process = subprocess.Popen(
player_command,
stdin=subprocess.PIPE,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
for audio_chunk in self.text_to_speak_stream(text):
player_process.stdin.write(audio_chunk)
player_process.stdin.flush()
player_process.stdin.close()
player_process.wait()
except Exception as e:
raise Exception(f"Audio player error: {e}")
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
return []
@@ -1,180 +0,0 @@
import os
import uuid
import json
import asyncio
import websockets
import ssl
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
from core.utils.util import parse_string_to_list
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
self.group_id = config.get("group_id")
self.api_key = config.get("api_key")
self.model = config.get("model")
# 初始化语音设置
default_voice_setting = {
"voice_id": "female-shaonv",
"speed": 1,
"vol": 1,
"pitch": 0,
"emotion": "happy",
}
default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]}
default_audio_setting = {
"sample_rate": 32000,
"bitrate": 128000,
"format": "mp3",
"channel": 1,
}
# 合并配置
self.voice_setting = {
**default_voice_setting,
**config.get("voice_setting", {}),
}
self.pronunciation_dict = {
**default_pronunciation_dict,
**config.get("pronunciation_dict", {}),
}
self.audio_setting = {
**default_audio_setting,
**config.get("audio_setting", {})
}
self.timber_weights = parse_string_to_list(config.get("timber_weights"))
# 设置语音ID
if config.get("private_voice"):
self.voice_setting["voice_id"] = config.get("private_voice")
elif config.get("voice_id"):
self.voice_setting["voice_id"] = config.get("voice_id")
# WebSocket配置
self.ws_url = "wss://api.minimaxi.com/ws/v1/t2a_v2"
self.headers = {
"Authorization": f"Bearer {self.api_key}",
"GroupId": self.group_id
}
self.audio_file_type = self.audio_setting.get("format", "mp3")
def generate_filename(self, extension=".mp3"):
"""生成唯一的音频文件名"""
return os.path.join(
self.output_file,
f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
)
async def _establish_connection(self):
"""建立WebSocket连接"""
ssl_context = ssl.create_default_context()
ssl_context.check_hostname = False
ssl_context.verify_mode = ssl.CERT_NONE
try:
ws = await websockets.connect(
self.ws_url,
additional_headers=self.headers,
ssl=ssl_context
)
connected = json.loads(await ws.recv())
if connected.get("event") == "connected_success":
print("连接成功")
return ws
return None
except Exception as e:
print(f"连接失败: {e}")
return None
async def _start_task(self, websocket):
"""发送任务开始请求"""
start_msg = {
"event": "task_start",
"model": self.model,
"voice_setting": self.voice_setting,
"pronunciation_dict": self.pronunciation_dict,
"audio_setting": self.audio_setting
}
if self.timber_weights and len(self.timber_weights) > 0:
start_msg["timber_weights"] = self.timber_weights
start_msg["voice_setting"]["voice_id"] = ""
await websocket.send(json.dumps(start_msg))
response = json.loads(await websocket.recv())
return response.get("event") == "task_started"
async def _continue_task(self, websocket, text):
"""发送继续请求并收集音频数据"""
await websocket.send(json.dumps({
"event": "task_continue",
"text": text
}))
audio_chunks = []
while True:
response = json.loads(await websocket.recv())
if "data" in response and "audio" in response["data"]:
audio_chunks.append(response["data"]["audio"])
if response.get("is_final"):
break
return "".join(audio_chunks)
async def _close_connection(self, websocket):
"""关闭连接"""
if websocket:
await websocket.send(json.dumps({"event": "task_finish"}))
await websocket.close()
print("连接已关闭")
async def text_to_speak(self, text, output_file=None):
"""主方法:文本转语音"""
ws = await self._establish_connection()
if not ws:
raise Exception("无法建立WebSocket连接")
try:
if not await self._start_task(ws):
raise Exception("任务启动失败")
hex_audio = await self._continue_task(ws, text)
audio_bytes = bytes.fromhex(hex_audio)
# 保存到文件或返回二进制数据
if output_file:
with open(output_file, "wb") as f:
f.write(audio_bytes)
print(f"音频已保存为{output_file}")
return output_file
else:
# 返回音频二进制数据(不播放)
return audio_bytes
finally:
await self._close_connection(ws)
async def main():
"""测试用主函数"""
# 示例配置
config = {
"group_id": "YOUR_GROUP_ID", # 替换为实际的group_id
"api_key": "YOUR_API_KEY", # 替换为实际的api_key
"model": "your-model", # 替换为实际的模型名称
"voice_id": "male-qn-qingse",
"voice_setting": {
"speed": 1.2,
"emotion": "happy"
}
}
tts = TTSProvider(config, delete_audio_file=True)
output_file = tts.generate_filename()
await tts.text_to_speak("这是一个测试文本,用于验证流式语音合成功能", output_file)
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,527 @@
import os
import time
import uuid
import json
import hmac
import queue
import base64
import hashlib
import asyncio
import traceback
import websockets
from asyncio import Task
from config.logger import setup_logging
from core.utils import opus_encoder_utils
from core.utils.tts import MarkdownCleaner
from urllib.parse import urlencode, urlparse
from core.providers.tts.base import TTSProviderBase
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
TAG = __name__
logger = setup_logging()
class XunfeiWSAuth:
@staticmethod
def create_auth_url(api_key, api_secret, api_url):
"""生成讯飞WebSocket认证URL"""
parsed_url = urlparse(api_url)
host = parsed_url.netloc
path = parsed_url.path
# 获取UTC时间,讯飞要求使用RFC1123格式
now = time.gmtime()
date = time.strftime('%a, %d %b %Y %H:%M:%S GMT', now)
# 构造签名字符串
signature_origin = f"host: {host}\ndate: {date}\nGET {path} HTTP/1.1"
# 计算签名
signature_sha = hmac.new(
api_secret.encode('utf-8'),
signature_origin.encode('utf-8'),
digestmod=hashlib.sha256
).digest()
signature_sha_base64 = base64.b64encode(signature_sha).decode(encoding='utf-8')
# 构造authorization
authorization_origin = f'api_key="{api_key}", algorithm="hmac-sha256", headers="host date request-line", signature="{signature_sha_base64}"'
authorization = base64.b64encode(authorization_origin.encode('utf-8')).decode(encoding='utf-8')
# 构造最终的WebSocket URL
v = {
"authorization": authorization,
"date": date,
"host": host
}
url = api_url + '?' + urlencode(v)
return url
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
# 设置为流式接口类型
self.interface_type = InterfaceType.DUAL_STREAM
# 基础配置
self.app_id = config.get("app_id")
self.api_key = config.get("api_key")
self.api_secret = config.get("api_secret")
# 接口地址
self.api_url = config.get("api_url", "wss://cbm01.cn-huabei-1.xf-yun.com/v1/private/mcd9m97e6")
# 音色配置
self.voice = config.get("voice", "x5_lingxiaoxuan_flow")
if config.get("private_voice"):
self.voice = config.get("private_voice")
# 音频参数配置
speed = config.get("speed", "50")
self.speed = int(speed) if speed else 50
volume = config.get("volume", "50")
self.volume = int(volume) if volume else 50
pitch = config.get("pitch", "50")
self.pitch = int(pitch) if pitch else 50
# 音频编码配置
self.format = config.get("format", "raw")
sample_rate = config.get("sample_rate", "24000")
self.sample_rate = int(sample_rate) if sample_rate else 24000
# 口语化配置
self.oral_level = config.get("oral_level", "mid")
spark_assist = config.get("spark_assist", "1")
self.spark_assist = int(spark_assist) if spark_assist else 1
stop_split = config.get("stop_split", "0")
self.stop_split = int(stop_split) if stop_split else 0
remain = config.get("remain", "0")
self.remain = int(remain) if remain else 0
# WebSocket配置
self.ws = None
self._monitor_task = None
# 序列号管理
self.text_seq = 0
# 创建Opus编码器
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
sample_rate=self.sample_rate, channels=1, frame_size_ms=60
)
# 验证必需参数
if not all([self.app_id, self.api_key, self.api_secret]):
raise ValueError("讯飞TTS需要配置app_id、api_key和api_secret")
async def _ensure_connection(self):
"""确保WebSocket连接可用"""
try:
logger.bind(tag=TAG).info("开始建立新连接...")
# 生成认证URL
auth_url = XunfeiWSAuth.create_auth_url(
self.api_key, self.api_secret, self.api_url
)
self.ws = await websockets.connect(
auth_url,
ping_interval=30,
ping_timeout=10,
close_timeout=10,
)
logger.bind(tag=TAG).info("WebSocket连接建立成功")
return self.ws
except Exception as e:
logger.bind(tag=TAG).error(f"建立连接失败: {str(e)}")
self.ws = None
raise
def tts_text_priority_thread(self):
"""流式文本处理线程"""
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} | 会话ID: {self.conn.sentence_id}"
)
if message.sentence_type == SentenceType.FIRST:
# 重置序列号
self.text_seq = 0
self.conn.client_abort = False
# 增加序列号
self.text_seq += 1
if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
continue
if message.sentence_type == SentenceType.FIRST:
# 初始化参数
try:
if not getattr(self.conn, "sentence_id", None):
self.conn.sentence_id = uuid.uuid4().hex
logger.bind(tag=TAG).info(f"自动生成新的 会话ID: {self.conn.sentence_id}")
logger.bind(tag=TAG).info("开始启动TTS会话...")
future = asyncio.run_coroutine_threadsafe(
self.start_session(self.conn.sentence_id),
loop=self.conn.loop,
)
future.result()
self.before_stop_play_files.clear()
logger.bind(tag=TAG).info("TTS会话启动成功")
except Exception as e:
logger.bind(tag=TAG).error(f"启动TTS会话失败: {str(e)}")
continue
# 处理文本内容
if ContentType.TEXT == message.content_type:
if message.content_detail:
try:
logger.bind(tag=TAG).debug(
f"开始发送TTS文本: {message.content_detail}"
)
future = asyncio.run_coroutine_threadsafe(
self.text_to_speak(message.content_detail, None),
loop=self.conn.loop,
)
future.result()
logger.bind(tag=TAG).debug("TTS文本发送成功")
except Exception as e:
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
# 不使用continue,确保后续处理不被中断
# 处理文件内容
if ContentType.FILE == message.content_type:
logger.bind(tag=TAG).info(
f"添加音频文件到待播放列表: {message.content_file}"
)
if message.content_file and os.path.exists(message.content_file):
# 先处理文件音频数据
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
# 处理会话结束
if message.sentence_type == SentenceType.LAST:
try:
logger.bind(tag=TAG).info("开始结束TTS会话...")
asyncio.run_coroutine_threadsafe(
self.finish_session(self.conn.sentence_id),
loop=self.conn.loop,
)
except Exception as e:
logger.bind(tag=TAG).error(f"结束TTS会话失败: {str(e)}")
continue
except queue.Empty:
continue
except Exception as e:
logger.bind(tag=TAG).error(
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
)
async def text_to_speak(self, text, _):
"""发送文本到TTS服务进行合成"""
try:
if self.ws is None:
logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本")
return
filtered_text = MarkdownCleaner.clean_markdown(text)
# 发送文本合成请求
run_request = self._build_base_request(status=1,text=filtered_text)
await self.ws.send(json.dumps(run_request))
return
except Exception as e:
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
if self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
raise
async def start_session(self, session_id):
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
try:
# 会话开始时检测上个会话的监听状态
if (
self._monitor_task is not None
and isinstance(self._monitor_task, Task)
and not self._monitor_task.done()
):
logger.bind(tag=TAG).info(
"检测到未完成的上个会话,关闭监听任务和连接..."
)
await self.close()
# 建立新连接
await self._ensure_connection()
# 启动监听任务
self._monitor_task = asyncio.create_task(self._start_monitor_tts_response())
# 发送会话启动请求
start_request = self._build_base_request(status=0)
await self.ws.send(json.dumps(start_request))
logger.bind(tag=TAG).info("会话启动请求已发送")
except Exception as e:
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
# 确保清理资源
await self.close()
raise
async def finish_session(self, session_id):
logger.bind(tag=TAG).info(f"关闭会话~~{session_id}")
try:
if self.ws:
# 发送会话结束请求
stop_request = self._build_base_request(status=2)
await self.ws.send(json.dumps(stop_request))
logger.bind(tag=TAG).info("会话结束请求已发送")
if self._monitor_task:
try:
await self._monitor_task
except Exception as e:
logger.bind(tag=TAG).error(f"等待监听任务完成时发生错误: {str(e)}")
finally:
self._monitor_task = None
except Exception as e:
logger.bind(tag=TAG).error(f"关闭会话失败: {str(e)}")
await self.close()
raise
async def close(self):
"""资源清理"""
if self._monitor_task:
try:
self._monitor_task.cancel()
await self._monitor_task
except asyncio.CancelledError:
pass
except Exception as e:
logger.bind(tag=TAG).warning(f"关闭时取消监听任务错误: {e}")
self._monitor_task = None
if self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
async def _start_monitor_tts_response(self):
"""监听TTS响应"""
try:
while not self.conn.stop_event.is_set():
try:
msg = await self.ws.recv()
# 检查客户端是否中止
if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止监听TTS响应")
break
try:
data = json.loads(msg)
header = data.get("header", {})
code = header.get("code")
if code == 0:
payload = data.get("payload", {})
audio_payload = payload.get("audio", {})
if audio_payload:
status = audio_payload.get("status", 0)
audio_data = audio_payload.get("audio", "")
if status == 0:
logger.bind(tag=TAG).debug("TTS合成已启动")
self.tts_audio_queue.put(
(SentenceType.FIRST, [], None)
)
elif status == 2:
logger.bind(tag=TAG).debug("收到结束状态的音频数据,TTS合成完成")
self._process_before_stop_play_files()
break
else:
if self.conn.tts_MessageText:
logger.bind(tag=TAG).info(
f"句子语音生成成功: {self.conn.tts_MessageText}"
)
self.tts_audio_queue.put(
(SentenceType.FIRST, [], self.conn.tts_MessageText)
)
self.conn.tts_MessageText = None
try:
audio_bytes = base64.b64decode(audio_data)
self.opus_encoder.encode_pcm_to_opus_stream(
audio_bytes, False, self.handle_opus
)
except Exception as e:
logger.bind(tag=TAG).error(f"处理音频数据失败: {e}")
else:
message = header.get("message", "未知错误")
logger.bind(tag=TAG).error(f"TTS合成错误: {code} - {message}")
break
except json.JSONDecodeError:
logger.bind(tag=TAG).warning("收到无效的JSON消息")
except websockets.ConnectionClosed:
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
break
except Exception as e:
logger.bind(tag=TAG).error(
f"处理TTS响应时出错: {e}\n{traceback.format_exc()}"
)
break
# 链接不可复用
if self.ws:
try:
await self.ws.close()
except:
pass
self.ws = None
# 监听任务退出时清理引用
finally:
self._monitor_task = None
def to_tts(self, text: str) -> list:
"""非流式TTS处理,用于测试及保存音频文件的场景"""
try:
# 创建新的事件循环
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
# 存储音频数据
audio_data = []
async def _generate_audio():
# 生成认证URL
auth_url = XunfeiWSAuth.create_auth_url(
self.api_key, self.api_secret, self.api_url
)
# 建立WebSocket连接
ws = await websockets.connect(
auth_url,
ping_interval=30,
ping_timeout=10,
close_timeout=10,
)
try:
filtered_text = MarkdownCleaner.clean_markdown(text)
text_request = self._build_base_request(status=2,text=filtered_text)
await ws.send(json.dumps(text_request))
task_finished = False
while not task_finished:
msg = await ws.recv()
data = json.loads(msg)
header = data.get("header", {})
code = header.get("code")
if code == 0:
payload = data.get("payload", {})
audio_payload = payload.get("audio", {})
if audio_payload:
status = audio_payload.get("status", 0)
audio_base64 = audio_payload.get("audio", "")
if status == 1:
try:
audio_bytes = base64.b64decode(audio_base64)
self.opus_encoder.encode_pcm_to_opus_stream(
audio_bytes,
end_of_stream=False,
callback=lambda opus: audio_data.append(opus)
)
except Exception as e:
logger.bind(tag=TAG).error(f"处理音频数据失败: {e}")
elif status == 2:
task_finished = True
logger.bind(tag=TAG).debug("TTS任务完成")
else:
message = header.get("message", "未知错误")
raise Exception(f"合成失败: {code} - {message}")
finally:
# 清理资源
try:
await ws.close()
except:
pass
loop.run_until_complete(_generate_audio())
loop.close()
return audio_data
except Exception as e:
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
return []
def _build_base_request(self, status,text=" "):
"""构建基础请求结构"""
return {
"header": {
"app_id": self.app_id,
"status": status,
},
"parameter": {
"oral": {
"oral_level": self.oral_level,
"spark_assist": self.spark_assist,
"stop_split": self.stop_split,
"remain": self.remain
},
"tts": {
"vcn": self.voice,
"speed": self.speed,
"volume": self.volume,
"pitch": self.pitch,
"bgs": 0,
"reg": 0,
"rdn": 0,
"rhy": 0,
"audio": {
"encoding": self.format,
"sample_rate": self.sample_rate,
"channels": 1,
"bit_depth": 16,
"frame_size": 0
}
}
},
"payload": {
"text": {
"encoding": "utf8",
"compress": "raw",
"format": "plain",
"status": status,
"seq": self.text_seq,
"text": base64.b64encode(text.encode('utf-8')).decode('utf-8')
}
}
}