x
This commit is contained in:
+458
-315
@@ -3,38 +3,73 @@ import json
|
||||
import websockets
|
||||
import numpy as np
|
||||
import sounddevice as sd
|
||||
from typing import Optional, Callable, Dict, Any, Coroutine, List
|
||||
from typing import Optional, Callable, Dict, Any, List
|
||||
from dataclasses import dataclass, field
|
||||
import uuid
|
||||
import copy
|
||||
from protocols import (
|
||||
EventType,
|
||||
MsgType,
|
||||
finish_connection,
|
||||
finish_session,
|
||||
receive_message,
|
||||
start_connection,
|
||||
start_session,
|
||||
task_request,
|
||||
wait_for_event,
|
||||
)
|
||||
|
||||
# ------------------------------
|
||||
# 配置常量(可根据需求调整)
|
||||
# ------------------------------
|
||||
DEFAULT_APPID = "7069844318"
|
||||
DEFAULT_ACCESS_TOKEN = "osFMEJr20SSTWRql43cJlZkAOg7iwvxu"
|
||||
DEFAULT_ENDPOINT = "wss://openspeech.bytedance.com/api/v3/tts/bidirection"
|
||||
DEFAULT_VOICE_TYPE = "zh_female_gaolengyujie_emo_v2_mars_bigtts"
|
||||
DEFAULT_ENCODING = "pcm"
|
||||
DEFAULT_SAMPLE_RATE = 16000
|
||||
|
||||
|
||||
@dataclass
|
||||
class TTSRequest:
|
||||
"""TTS请求对象(带唯一标识)"""
|
||||
tts_text: str
|
||||
mode: str = "预训练音色"
|
||||
sft_spk: str = ""
|
||||
seed: int = field(default_factory=lambda: np.random.randint(1, 100000000))
|
||||
stream: bool = True
|
||||
speed: float = 1.0
|
||||
prompt_wav: str = ""
|
||||
instruct_text: str = ""
|
||||
voice_type: str = DEFAULT_VOICE_TYPE
|
||||
encoding: str = DEFAULT_ENCODING
|
||||
speed: float = 1.0 # 语速(字节跳动TTS支持,需服务端兼容)
|
||||
stream: bool = True # 是否流式合成
|
||||
request_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 唯一请求ID
|
||||
session_id: str = field(default_factory=lambda: str(uuid.uuid4())) # 会话ID(每个请求一个会话)
|
||||
|
||||
|
||||
class CosyVoiceTTSSocketClient:
|
||||
"""CosyVoice TTS WebSocket 客户端(异步/流式/带任务队列)"""
|
||||
class ByteDanceTTSSocketClient:
|
||||
"""字节跳动 TTS WebSocket 客户端(异步/流式/带任务队列)"""
|
||||
|
||||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts", max_queue_size: int = 100):
|
||||
def __init__(
|
||||
self,
|
||||
appid: str = DEFAULT_APPID,
|
||||
access_token: str = DEFAULT_ACCESS_TOKEN,
|
||||
endpoint: str = DEFAULT_ENDPOINT,
|
||||
max_queue_size: int = 100
|
||||
):
|
||||
"""
|
||||
初始化客户端
|
||||
:param ws_url: WebSocket 服务端地址
|
||||
:param appid: 字节跳动APP ID
|
||||
:param access_token: 访问令牌
|
||||
:param endpoint: WebSocket 服务端地址
|
||||
:param max_queue_size: 最大队列长度(防止内存溢出)
|
||||
"""
|
||||
self.ws_url = ws_url
|
||||
# 基础配置
|
||||
self.appid = appid
|
||||
self.access_token = access_token
|
||||
self.endpoint = endpoint
|
||||
self.max_queue_size = max_queue_size
|
||||
|
||||
# WebSocket 连接状态
|
||||
self.websocket: Optional[websockets.WebSocketClientProtocol] = None
|
||||
self.is_connected = False
|
||||
self.is_processing = False # 是否正在处理请求
|
||||
self.logid: Optional[str] = None # 服务端返回的日志ID
|
||||
|
||||
# 异步任务队列(FIFO)
|
||||
self.request_queue: asyncio.Queue[TTSRequest] = asyncio.Queue(maxsize=max_queue_size)
|
||||
@@ -42,45 +77,84 @@ class CosyVoiceTTSSocketClient:
|
||||
# 回调函数定义(所有回调都带request_id,方便关联请求)
|
||||
self.on_task_enqueue: Callable[[str], None] = lambda req_id: None # 任务入队回调
|
||||
self.on_start: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成开始回调
|
||||
self.on_audio_chunk: Callable[[str, np.ndarray], None] = lambda req_id, chunk: None # 音频块回调
|
||||
self.on_audio_chunk: Callable[[str, bytes], None] = lambda req_id, chunk: None # 音频块回调(原始字节)
|
||||
self.on_end: Callable[[str, Dict[str, Any]], None] = lambda req_id, data: None # 合成结束回调
|
||||
self.on_error: Callable[[str, str], None] = lambda req_id, msg: None # 错误回调
|
||||
self.on_queue_full: Callable[[str], None] = lambda req_id: None # 队列满回调
|
||||
|
||||
# 音频播放流
|
||||
# 音频播放相关(支持MP3格式直接播放)
|
||||
self.play_stream: Optional[sd.OutputStream] = None
|
||||
self.current_req_id: Optional[str] = None
|
||||
self.enable_playback: bool = True # 是否启用实时播放
|
||||
|
||||
# 是否启用播放功能
|
||||
self.enable_playback: bool = True
|
||||
|
||||
# 外部回调函数,用于返回完整结果
|
||||
# 外部回调函数(返回完整结果)
|
||||
self.external_callback: Optional[Callable[[str, Dict[str, Any]], None]] = None
|
||||
|
||||
# 存储每个请求的音频数据
|
||||
self.audio_buffers: Dict[str, List[np.ndarray]] = {}
|
||||
# 存储每个请求的完整音频数据(原始字节)
|
||||
self.audio_buffers: Dict[str, List[bytes]] = {}
|
||||
|
||||
def _get_resource_id(self, voice_type: str) -> str:
|
||||
"""根据音色类型获取资源ID(字节跳动TTS协议要求)"""
|
||||
if voice_type.startswith("S_"):
|
||||
return "volc.megatts.default"
|
||||
return "volc.service_type.10029"
|
||||
|
||||
async def _create_websocket_connection(self):
|
||||
"""创建WebSocket连接(内部使用)"""
|
||||
headers = {
|
||||
"X-Api-App-Key": self.appid,
|
||||
"X-Api-Access-Key": self.access_token,
|
||||
"X-Api-Resource-Id": self._get_resource_id(DEFAULT_VOICE_TYPE), # 用默认音色获取资源ID
|
||||
"X-Api-Connect-Id": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
print(f"连接到 TTS 服务端: {self.endpoint}")
|
||||
self.websocket = await websockets.connect(
|
||||
self.endpoint,
|
||||
additional_headers=headers,
|
||||
max_size=10 * 1024 * 1024 # 10MB缓冲区
|
||||
)
|
||||
self.is_connected = True
|
||||
self.logid = self.websocket.response.headers.get("x-tt-logid")
|
||||
print(f"连接成功,LogID: {self.logid}")
|
||||
|
||||
# 发送连接启动指令
|
||||
await start_connection(self.websocket)
|
||||
await wait_for_event(
|
||||
self.websocket, MsgType.FullServerResponse, EventType.ConnectionStarted
|
||||
)
|
||||
print("TTS连接已初始化完成")
|
||||
|
||||
async def connect(self):
|
||||
"""建立 WebSocket 连接(初始化一次)"""
|
||||
"""建立WebSocket连接(外部调用,初始化一次)"""
|
||||
if not self.is_connected:
|
||||
try:
|
||||
self.websocket = await websockets.connect(self.ws_url)
|
||||
self.is_connected = True
|
||||
print(f"成功连接到 TTS 服务端: {self.ws_url}")
|
||||
|
||||
await self._create_websocket_connection()
|
||||
# 启动队列消费协程(后台运行)
|
||||
asyncio.create_task(self._consume_queue())
|
||||
except Exception as e:
|
||||
raise ConnectionError(f"连接失败: {str(e)}")
|
||||
error_msg = f"连接失败: {str(e)}"
|
||||
print(error_msg)
|
||||
raise ConnectionError(error_msg)
|
||||
|
||||
async def disconnect(self):
|
||||
"""关闭 WebSocket 连接"""
|
||||
"""关闭WebSocket连接"""
|
||||
if self.is_connected and self.websocket:
|
||||
await self.websocket.close()
|
||||
self.is_connected = False
|
||||
self.websocket = None
|
||||
print("已断开与 TTS 服务端的连接")
|
||||
try:
|
||||
# 发送连接结束指令
|
||||
await finish_connection(self.websocket)
|
||||
await wait_for_event(
|
||||
self.websocket, MsgType.FullServerResponse, EventType.ConnectionFinished
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"关闭连接时异常: {str(e)}")
|
||||
finally:
|
||||
await self.websocket.close()
|
||||
self.is_connected = False
|
||||
self.websocket = None
|
||||
print("已断开与TTS服务端的连接")
|
||||
|
||||
# 清理音频播放流
|
||||
if self.play_stream:
|
||||
self.play_stream.stop()
|
||||
self.play_stream.close()
|
||||
@@ -91,281 +165,377 @@ class CosyVoiceTTSSocketClient:
|
||||
self.external_callback = callback
|
||||
|
||||
def set_playback_enabled(self, enabled: bool):
|
||||
"""设置是否启用音频播放"""
|
||||
"""设置是否启用音频实时播放"""
|
||||
self.enable_playback = enabled
|
||||
print(f"音频实时播放已{'启用' if enabled else '禁用'}")
|
||||
|
||||
async def add_tts_request(self, tts_text: str, **kwargs) -> str:
|
||||
async def synthesize(self, tts_text: str, **kwargs) -> str:
|
||||
"""
|
||||
添加TTS请求到队列(异步非阻塞)
|
||||
异步非阻塞添加TTS请求到队列
|
||||
:param tts_text: 要合成的文本
|
||||
:param kwargs: 其他TTS参数
|
||||
:param kwargs: 其他TTS参数(voice_type, encoding, speed等)
|
||||
:return: 唯一请求ID
|
||||
"""
|
||||
# 创建请求对象
|
||||
# 创建请求对象(支持覆盖默认参数)
|
||||
request = TTSRequest(tts_text=tts_text, **kwargs)
|
||||
req_id = request.request_id
|
||||
|
||||
# 初始化音频缓冲区
|
||||
self.audio_buffers[req_id] = []
|
||||
|
||||
# 异步入队
|
||||
# 异步入队(非阻塞)
|
||||
try:
|
||||
await self.request_queue.put(request)
|
||||
self.on_task_enqueue(req_id)
|
||||
print(f"请求 {req_id} 已加入队列,当前队列长度: {self.request_queue.qsize()}")
|
||||
print(f"请求 [{req_id[:8]}] 已加入队列,当前队列长度: {self.request_queue.qsize()}")
|
||||
return req_id
|
||||
except asyncio.QueueFull:
|
||||
self.on_queue_full(req_id)
|
||||
raise Exception(f"队列已满,请求 {req_id} 入队失败")
|
||||
error_msg = f"队列已满(最大长度{self.max_queue_size}),请求 [{req_id[:8]}] 入队失败"
|
||||
print(error_msg)
|
||||
raise Exception(error_msg)
|
||||
|
||||
async def _consume_queue(self):
|
||||
"""消费队列(后台协程,自动处理排队请求)"""
|
||||
print("队列消费协程已启动")
|
||||
while True:
|
||||
try:
|
||||
# 等待队列中有请求(阻塞)
|
||||
# 等待队列中有请求(阻塞,直到有任务)
|
||||
request = await self.request_queue.get()
|
||||
req_id = request.request_id
|
||||
|
||||
# 标记为处理中
|
||||
self.is_processing = True
|
||||
print(f"开始处理请求 {req_id},剩余队列长度: {self.request_queue.qsize()}")
|
||||
# print(f"\n开始处理请求 [{req_id[:8]}],剩余队列长度: {self.request_queue.qsize()}")
|
||||
|
||||
# 处理当前请求
|
||||
# 处理单个请求
|
||||
await self._process_single_request(request)
|
||||
|
||||
# 标记任务完成
|
||||
# 标记任务完成(让Queue知道可以继续)
|
||||
self.request_queue.task_done()
|
||||
self.is_processing = False
|
||||
|
||||
except Exception as e:
|
||||
print(f"队列消费异常: {str(e)}")
|
||||
error_msg = f"队列消费异常: {str(e)}"
|
||||
print(error_msg)
|
||||
self.is_processing = False
|
||||
# 短暂等待后继续消费,避免死循环
|
||||
# 短暂等待,避免死循环占用CPU
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
async def _process_single_request(self, request: TTSRequest):
|
||||
"""处理单个TTS请求"""
|
||||
req_id = request.request_id
|
||||
|
||||
# 参数校验
|
||||
if not request.tts_text:
|
||||
self.on_error(req_id, "合成文本不能为空")
|
||||
return
|
||||
|
||||
# 确保已连接
|
||||
if not self.is_connected:
|
||||
try:
|
||||
await self.connect()
|
||||
except Exception as e:
|
||||
self.on_error(req_id, f"连接失败: {str(e)}")
|
||||
return
|
||||
|
||||
# 构造请求数据
|
||||
request_data = {
|
||||
"tts_text": request.tts_text,
|
||||
"mode": request.mode,
|
||||
"sft_spk": request.sft_spk,
|
||||
"seed": request.seed,
|
||||
"stream": request.stream,
|
||||
"speed": request.speed,
|
||||
"prompt_wav": request.prompt_wav,
|
||||
"instruct_text": request.instruct_text
|
||||
def _build_base_request(self, request: TTSRequest) -> Dict[str, Any]:
|
||||
"""构建字节跳动TTS基础请求参数"""
|
||||
aaa = {
|
||||
"user": {"uid": str(uuid.uuid4())},
|
||||
"namespace": "BidirectionalTTS",
|
||||
"req_params": {
|
||||
"speaker": request.voice_type,
|
||||
"audio_params": {
|
||||
"format": request.encoding,
|
||||
"sample_rate": DEFAULT_SAMPLE_RATE,
|
||||
"enable_timestamp": True,
|
||||
},
|
||||
"additions": json.dumps({"disable_markdown_filter": False}),
|
||||
"speed": request.speed, # 语速参数(需服务端支持)
|
||||
},
|
||||
}
|
||||
print('aaa', aaa)
|
||||
return aaa
|
||||
async def _send_text_stream(self, request: TTSRequest, session_id: str):
|
||||
"""流式发送文本(逐字符发送,字节跳动TTS流式协议要求)"""
|
||||
base_request = self._build_base_request(request)
|
||||
text = request.tts_text.strip()
|
||||
|
||||
try:
|
||||
# 发送请求
|
||||
await self.websocket.send(json.dumps(request_data))
|
||||
|
||||
# 处理返回的流式数据
|
||||
await self._handle_response(req_id)
|
||||
|
||||
except websockets.exceptions.ConnectionClosed:
|
||||
self.is_connected = False
|
||||
self.on_error(req_id, "连接已关闭")
|
||||
except Exception as e:
|
||||
self.on_error(req_id, f"处理请求失败: {str(e)}")
|
||||
|
||||
async def _handle_response(self, req_id: str):
|
||||
"""处理单个请求的服务端响应"""
|
||||
if not self.websocket:
|
||||
if not text:
|
||||
print(f"请求 [{request.request_id[:8]}] 文本为空,跳过发送")
|
||||
return
|
||||
|
||||
# 逐字符发送(控制发送速率,避免拥塞)
|
||||
for char in text:
|
||||
if not self.is_connected or not self.websocket:
|
||||
raise ConnectionError("连接已断开,无法继续发送文本")
|
||||
|
||||
# 构建单个字符的任务请求
|
||||
task_req = copy.deepcopy(base_request)
|
||||
task_req["event"] = EventType.TaskRequest
|
||||
task_req["req_params"]["text"] = char
|
||||
|
||||
# 发送任务请求
|
||||
await task_request(
|
||||
self.websocket,
|
||||
json.dumps(task_req).encode("utf-8"),
|
||||
session_id
|
||||
)
|
||||
# 控制发送速率(5ms/字符,可调整)
|
||||
await asyncio.sleep(0.005)
|
||||
|
||||
# 发送会话结束指令
|
||||
await finish_session(self.websocket, session_id)
|
||||
print(f"请求 [{request.request_id[:8]}] 文本发送完成")
|
||||
|
||||
async def _handle_audio_response(self, req_id: str, session_id: str, request: TTSRequest) -> Dict[str, Any]:
|
||||
"""处理服务端的流式音频响应"""
|
||||
if not self.websocket:
|
||||
raise ConnectionError("WebSocket连接未建立")
|
||||
|
||||
sample_rate = DEFAULT_SAMPLE_RATE
|
||||
audio_received = False
|
||||
|
||||
try:
|
||||
sample_rate = None
|
||||
while True:
|
||||
# 接收服务端消息(异步)
|
||||
response = await self.websocket.recv()
|
||||
data = json.loads(response)
|
||||
# 接收服务端消息(异步阻塞)
|
||||
msg = await receive_message(self.websocket)
|
||||
|
||||
# 根据状态分发到不同回调(都带req_id)
|
||||
status = data.get("status")
|
||||
if status == "start":
|
||||
# 合成开始 - 返回采样率等信息
|
||||
sample_rate = data.get("sample_rate")
|
||||
self.on_start(req_id, data)
|
||||
if msg.type == MsgType.FullServerResponse:
|
||||
# 完整响应(开始/结束/错误)
|
||||
if msg.event == EventType.SessionStarted:
|
||||
# 会话开始回调
|
||||
start_data = {
|
||||
"session_id": session_id,
|
||||
"sample_rate": sample_rate,
|
||||
"encoding": request.encoding,
|
||||
"voice_type": request.voice_type,
|
||||
"logid": self.logid
|
||||
}
|
||||
self.on_start(req_id, start_data)
|
||||
print(f"请求 [{req_id[:8]}] 合成开始")
|
||||
|
||||
elif status == "stream":
|
||||
# 流式音频块 - 转为 numpy 数组
|
||||
audio_chunk = np.array(data["audio_chunk"], dtype=np.float32)
|
||||
elif msg.event == EventType.SessionFinished:
|
||||
# 会话结束,退出循环
|
||||
end_data = {"session_id": session_id, "message": "合成完成"}
|
||||
self.on_end(req_id, end_data)
|
||||
print(f"请求 [{req_id[:8]}] 合成结束")
|
||||
break
|
||||
|
||||
# 保存到缓冲区
|
||||
self.audio_buffers[req_id].append(audio_chunk)
|
||||
elif msg.type == MsgType.AudioOnlyServer:
|
||||
# 流式音频数据(原始字节)
|
||||
audio_chunk = msg.payload
|
||||
if audio_chunk:
|
||||
audio_received = True
|
||||
# 保存到缓冲区
|
||||
self.audio_buffers[req_id].append(audio_chunk)
|
||||
# 音频块回调
|
||||
self.on_audio_chunk(req_id, audio_chunk)
|
||||
# 实时播放(如果启用)
|
||||
await self._play_audio_chunk(req_id, audio_chunk, sample_rate)
|
||||
|
||||
# 音频块回调
|
||||
self.on_audio_chunk(req_id, audio_chunk)
|
||||
else:
|
||||
# 未知消息类型
|
||||
raise RuntimeError(f"收到未知消息类型: {msg.type}, 内容: {msg}")
|
||||
|
||||
elif status == "end":
|
||||
# 合成结束
|
||||
self.on_end(req_id, data)
|
||||
|
||||
# 组装完整音频数据
|
||||
full_audio = np.concatenate(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else np.array(
|
||||
[], dtype=np.float32)
|
||||
|
||||
# 准备返回给外部回调的数据
|
||||
result_data = {
|
||||
"status": "completed",
|
||||
"request_id": req_id,
|
||||
"sample_rate": sample_rate,
|
||||
"audio_data": full_audio,
|
||||
"audio_length": len(full_audio),
|
||||
"message": data.get("msg", "合成完成")
|
||||
}
|
||||
|
||||
# 调用外部回调(非阻塞)
|
||||
if self.external_callback:
|
||||
# 使用create_task确保回调不会阻塞主流程
|
||||
asyncio.create_task(self._call_external_callback(req_id, result_data))
|
||||
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
break
|
||||
|
||||
elif status == "error":
|
||||
# 错误处理
|
||||
error_msg = data.get("msg", "未知错误")
|
||||
self.on_error(req_id, error_msg)
|
||||
|
||||
# 准备错误结果数据
|
||||
error_data = {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"message": error_msg
|
||||
}
|
||||
|
||||
# 调用外部回调
|
||||
if self.external_callback:
|
||||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||||
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
break
|
||||
# 组装完整结果
|
||||
full_audio = b"".join(self.audio_buffers[req_id]) if self.audio_buffers[req_id] else b""
|
||||
return {
|
||||
"status": "completed",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"sample_rate": sample_rate,
|
||||
"encoding": request.encoding,
|
||||
"audio_data": full_audio, # 完整音频字节数据
|
||||
"audio_length": len(full_audio),
|
||||
"message": "合成成功" if audio_received else "合成完成但未收到音频数据"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"接收数据异常: {str(e)}"
|
||||
error_msg = f"处理音频响应异常: {str(e)}"
|
||||
self.on_error(req_id, error_msg)
|
||||
|
||||
# 准备错误结果数据
|
||||
error_data = {
|
||||
return {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
}
|
||||
|
||||
# 调用外部回调
|
||||
if self.external_callback:
|
||||
asyncio.create_task(self._call_external_callback(req_id, error_data))
|
||||
async def _play_audio_chunk(self, req_id: str, chunk: bytes, sample_rate: int):
|
||||
"""实时播放音频块(支持MP3格式)"""
|
||||
if not self.enable_playback:
|
||||
return
|
||||
|
||||
# 确保当前请求是正在播放的请求
|
||||
if self.current_req_id is None:
|
||||
self.current_req_id = req_id
|
||||
|
||||
if req_id != self.current_req_id:
|
||||
# 切换请求时,重置播放流
|
||||
if self.play_stream:
|
||||
self.play_stream.stop()
|
||||
self.play_stream.close()
|
||||
self.current_req_id = req_id
|
||||
|
||||
try:
|
||||
# 初始化播放流(如果未初始化)
|
||||
if not self.play_stream:
|
||||
self.play_stream = sd.OutputStream(
|
||||
samplerate=sample_rate,
|
||||
channels=1, # 单声道
|
||||
dtype=np.float32
|
||||
)
|
||||
self.play_stream.start()
|
||||
|
||||
# MP3字节 → 音频数组(直接播放)
|
||||
# 注意:sounddevice默认支持PCM格式,如果是MP3需要解码,这里简化处理(实际使用建议用pydub解码)
|
||||
# 如需支持MP3播放,请安装 pydub: pip install pydub ffmpeg
|
||||
try:
|
||||
# 简化处理:假设服务端返回PCM(如果是MP3,需替换为解码逻辑)
|
||||
audio_array = np.frombuffer(chunk, dtype=np.float32)
|
||||
if audio_array.size > 0:
|
||||
self.play_stream.write(audio_array)
|
||||
except Exception as e:
|
||||
print(f"音频播放异常: {str(e)},请确保音频格式正确")
|
||||
|
||||
except Exception as e:
|
||||
print(f"播放流初始化失败: {str(e)}")
|
||||
|
||||
async def _process_single_request(self, request: TTSRequest):
|
||||
"""处理单个TTS请求(完整流程:连接→启动会话→流式发送文本→接收音频→回调结果)"""
|
||||
req_id = request.request_id
|
||||
session_id = request.session_id
|
||||
|
||||
# 参数校验
|
||||
if not request.tts_text.strip():
|
||||
error_msg = "合成文本不能为空"
|
||||
self.on_error(req_id, error_msg)
|
||||
self._send_external_callback(req_id, {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
})
|
||||
return
|
||||
|
||||
# 确保连接已建立(断开时自动重连)
|
||||
if not self.is_connected:
|
||||
print(f"请求 [{req_id[:8]}] 处理时连接已断开,尝试重连...")
|
||||
try:
|
||||
await self._create_websocket_connection()
|
||||
except Exception as e:
|
||||
error_msg = f"重连失败: {str(e)}"
|
||||
self.on_error(req_id, error_msg)
|
||||
self._send_external_callback(req_id, {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
})
|
||||
return
|
||||
|
||||
try:
|
||||
# 1. 启动会话
|
||||
base_request = self._build_base_request(request)
|
||||
start_session_req = copy.deepcopy(base_request)
|
||||
start_session_req["event"] = EventType.StartSession
|
||||
|
||||
await start_session(
|
||||
self.websocket,
|
||||
json.dumps(start_session_req).encode("utf-8"),
|
||||
session_id
|
||||
)
|
||||
# 等待会话启动成功
|
||||
await wait_for_event(
|
||||
self.websocket, MsgType.FullServerResponse, EventType.SessionStarted
|
||||
)
|
||||
|
||||
# 2. 异步流式发送文本(后台任务,不阻塞接收音频)
|
||||
send_task = asyncio.create_task(self._send_text_stream(request, session_id))
|
||||
|
||||
# 3. 接收并处理音频响应
|
||||
result_data = await self._handle_audio_response(req_id, session_id, request)
|
||||
|
||||
# 4. 等待文本发送任务完成
|
||||
await send_task
|
||||
|
||||
# 5. 发送外部回调
|
||||
self._send_external_callback(req_id, result_data)
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"处理请求 [{req_id[:8]}] 异常: {str(e)}"
|
||||
self.on_error(req_id, error_msg)
|
||||
self._send_external_callback(req_id, {
|
||||
"status": "error",
|
||||
"request_id": req_id,
|
||||
"session_id": session_id,
|
||||
"message": error_msg
|
||||
})
|
||||
finally:
|
||||
# 清理缓冲区
|
||||
if req_id in self.audio_buffers:
|
||||
del self.audio_buffers[req_id]
|
||||
|
||||
async def _call_external_callback(self, req_id: str, result_data: Dict[str, Any]):
|
||||
"""异步调用外部回调函数"""
|
||||
try:
|
||||
if self.external_callback:
|
||||
# 确保回调是异步的,不会阻塞
|
||||
if asyncio.iscoroutinefunction(self.external_callback):
|
||||
await self.external_callback(req_id, result_data)
|
||||
else:
|
||||
# 如果是同步函数,在线程池中执行
|
||||
await asyncio.get_event_loop().run_in_executor(
|
||||
None, self.external_callback, req_id, result_data
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"外部回调执行异常: {str(e)}")
|
||||
|
||||
def _play_audio_chunk(self, chunk: np.ndarray):
|
||||
"""播放音频块(非阻塞)"""
|
||||
if not self.enable_playback or not self.play_stream:
|
||||
def _send_external_callback(self, req_id: str, result_data: Dict[str, Any]):
|
||||
"""发送外部回调(支持同步/异步回调函数)"""
|
||||
if not self.external_callback:
|
||||
return
|
||||
|
||||
try:
|
||||
if chunk.size > 0:
|
||||
self.play_stream.write(chunk)
|
||||
# 异步回调:直接await
|
||||
if asyncio.iscoroutinefunction(self.external_callback):
|
||||
asyncio.create_task(self.external_callback(req_id, result_data))
|
||||
# 同步回调:在线程池中执行(避免阻塞事件循环)
|
||||
else:
|
||||
asyncio.get_event_loop().run_in_executor(
|
||||
None, self.external_callback, req_id, result_data
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"音频播放异常: {str(e)}")
|
||||
print(f"外部回调执行异常: {str(e)}")
|
||||
|
||||
async def wait_queue_empty(self):
|
||||
"""等待队列所有任务处理完成(阻塞)"""
|
||||
async def wait_all_completed(self):
|
||||
"""等待队列中所有任务处理完成(阻塞)"""
|
||||
await self.request_queue.join()
|
||||
print("所有队列任务已处理完成")
|
||||
print("\n所有队列任务已处理完成")
|
||||
|
||||
|
||||
# ------------------------------
|
||||
# 使用示例
|
||||
# 使用示例(与你提供的风格完全一致)
|
||||
# ------------------------------
|
||||
class TTSManager:
|
||||
"""TTS管理器 - 供外部代码调用"""
|
||||
"""TTS管理器 - 供外部代码调用(封装客户端,简化使用)"""
|
||||
|
||||
def __init__(self, ws_url: str = "ws://localhost:50000/ws/tts"):
|
||||
self.client = CosyVoiceTTSSocketClient(ws_url)
|
||||
self._setup_callbacks()
|
||||
def __init__(
|
||||
self,
|
||||
appid: str = DEFAULT_APPID,
|
||||
access_token: str = DEFAULT_ACCESS_TOKEN,
|
||||
endpoint: str = DEFAULT_ENDPOINT
|
||||
):
|
||||
self.client = ByteDanceTTSSocketClient(
|
||||
appid=appid,
|
||||
access_token=access_token,
|
||||
endpoint=endpoint
|
||||
)
|
||||
self._setup_internal_callbacks()
|
||||
|
||||
def _setup_callbacks(self):
|
||||
"""设置内部回调"""
|
||||
def _setup_internal_callbacks(self):
|
||||
"""设置内部回调(日志/状态提示)"""
|
||||
|
||||
# 合成开始回调
|
||||
def on_tts_start(req_id, data):
|
||||
print(f"\n🎤 开始合成 [{req_id[:8]}] - 采样率: {data.get('sample_rate')}")
|
||||
def on_task_enqueue(req_id: str):
|
||||
"""任务入队回调"""
|
||||
print(f"📥 任务 [{req_id[:8]}] 已入队")
|
||||
|
||||
# 初始化播放流(如果需要播放)
|
||||
if self.client.enable_playback:
|
||||
sample_rate = data.get("sample_rate", 24000)
|
||||
if self.client.play_stream:
|
||||
self.client.play_stream.stop()
|
||||
self.client.play_stream.close()
|
||||
def on_tts_start(req_id: str, data: Dict[str, Any]):
|
||||
"""合成开始回调"""
|
||||
print(f"🎤 合成开始 [{req_id[:8]}] - 采样率: {data['sample_rate']}, 编码: {data['encoding']}")
|
||||
|
||||
self.client.play_stream = sd.OutputStream(
|
||||
samplerate=sample_rate,
|
||||
channels=1,
|
||||
dtype=np.float32
|
||||
)
|
||||
self.client.play_stream.start()
|
||||
self.client.current_req_id = req_id
|
||||
def on_audio_chunk(req_id: str, chunk: bytes):
|
||||
"""音频块回调(内部仅打印日志,外部通过external_callback获取)"""
|
||||
print(f"🔊 收到音频块 [{req_id[:8]}] - 大小: {len(chunk)}字节", end="\r")
|
||||
|
||||
# 音频块回调
|
||||
def on_audio_chunk(req_id, chunk):
|
||||
if req_id == self.client.current_req_id:
|
||||
self.client._play_audio_chunk(chunk)
|
||||
def on_tts_end(req_id: str, data: Dict[str, Any]):
|
||||
"""合成结束回调"""
|
||||
print(f"\n🏁 合成结束 [{req_id[:8]}] - 会话ID: {data['session_id']}")
|
||||
|
||||
# 合成结束回调
|
||||
def on_tts_end(req_id, data):
|
||||
print(f"🏁 合成完成 [{req_id[:8]}] - {data.get('msg', '完成')}")
|
||||
def on_tts_error(req_id: str, msg: str):
|
||||
"""错误回调"""
|
||||
print(f"\n❌ 合成失败 [{req_id[:8]}] - 错误: {msg}")
|
||||
|
||||
# 错误回调
|
||||
def on_tts_error(req_id, msg):
|
||||
print(f"❌ 合成失败 [{req_id[:8]}] - {msg}")
|
||||
def on_queue_full(req_id: str):
|
||||
"""队列满回调"""
|
||||
print(f"⚠️ 队列已满,请求 [{req_id[:8]}] 入队失败")
|
||||
|
||||
# 绑定内部回调
|
||||
self.client.on_task_enqueue = on_task_enqueue
|
||||
self.client.on_start = on_tts_start
|
||||
self.client.on_audio_chunk = on_audio_chunk
|
||||
self.client.on_end = on_tts_end
|
||||
self.client.on_error = on_tts_error
|
||||
self.client.on_queue_full = on_queue_full
|
||||
|
||||
async def initialize(self):
|
||||
"""初始化连接"""
|
||||
@@ -376,25 +546,25 @@ class TTSManager:
|
||||
await self.client.disconnect()
|
||||
|
||||
def set_result_callback(self, callback: Callable[[str, Dict[str, Any]], None]):
|
||||
"""设置结果回调函数"""
|
||||
"""设置外部结果回调(获取完整音频数据)"""
|
||||
self.client.set_external_callback(callback)
|
||||
|
||||
def set_playback_enabled(self, enabled: bool):
|
||||
"""设置是否启用播放"""
|
||||
"""设置是否启用实时播放"""
|
||||
self.client.set_playback_enabled(enabled)
|
||||
|
||||
async def synthesize(self, text: str, **kwargs) -> str:
|
||||
"""
|
||||
异步合成文本(非阻塞)
|
||||
异步非阻塞合成文本
|
||||
:param text: 要合成的文本
|
||||
:param kwargs: 其他参数(mode, sft_spk, speed等)
|
||||
:param kwargs: 其他参数(voice_type, encoding, speed等)
|
||||
:return: 请求ID
|
||||
"""
|
||||
return await self.client.add_tts_request(text, **kwargs)
|
||||
return await self.client.synthesize(text, **kwargs)
|
||||
|
||||
async def wait_all_completed(self):
|
||||
"""等待所有任务完成"""
|
||||
await self.client.wait_queue_empty()
|
||||
await self.client.wait_all_completed()
|
||||
|
||||
|
||||
# ------------------------------
|
||||
@@ -402,147 +572,120 @@ class TTSManager:
|
||||
# ------------------------------
|
||||
async def external_usage_example():
|
||||
"""外部代码使用示例"""
|
||||
# 1. 创建TTS管理器(可替换为自己的appid和access_token)
|
||||
tts_manager = TTSManager(
|
||||
appid=DEFAULT_APPID,
|
||||
access_token=DEFAULT_ACCESS_TOKEN,
|
||||
endpoint=DEFAULT_ENDPOINT
|
||||
)
|
||||
|
||||
# 1. 创建TTS管理器
|
||||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
|
||||
# 2. 设置结果回调函数(接收完整结果)
|
||||
def handle_tts_result1(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS结果回调"""
|
||||
# 2. 设置外部结果回调(获取完整音频数据)
|
||||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS完整结果(同步回调)"""
|
||||
status = result.get("status")
|
||||
|
||||
if status == "completed":
|
||||
audio_data = result.get("audio_data")
|
||||
sample_rate = result.get("sample_rate")
|
||||
encoding = result.get("encoding")
|
||||
audio_length = result.get("audio_length")
|
||||
|
||||
print(f"✅ 收到TTS结果 [{req_id[:8]}]: 长度{audio_length}采样点, 采样率{sample_rate}Hz")
|
||||
print(f"\n✅ 收到完整结果 [{req_id[:8]}] - 长度: {audio_length}字节, 编码: {encoding}")
|
||||
|
||||
# 这里可以保存音频文件或进行其他处理
|
||||
# 注意:audio_data是完整的numpy数组
|
||||
# 保存音频文件
|
||||
filename = f"tts_output_{req_id[:8]}.{encoding}"
|
||||
with open(filename, "wb") as f:
|
||||
f.write(audio_data)
|
||||
print(f"💾 音频文件已保存: {filename}")
|
||||
|
||||
elif status == "error":
|
||||
error_msg = result.get("message")
|
||||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
print(f"\n❌ 请求 [{req_id[:8]}] 处理失败: {error_msg}")
|
||||
|
||||
def handle_tts_result(req_id: str, result: Dict[str, Any]):
|
||||
"""处理TTS结果回调 - 保存为PCM文件"""
|
||||
print('处理TTS结果回调', result)
|
||||
status = result.get("status")
|
||||
print('status')
|
||||
if status == "completed":
|
||||
audio_data = result.get("audio_data")
|
||||
sample_rate = result.get("sample_rate")
|
||||
|
||||
if audio_data is not None and len(audio_data) > 0:
|
||||
# 转换为PCM数据
|
||||
pcm_data = (audio_data.astype(np.float32) * 32767).astype(np.int16)
|
||||
pcm_bytes = pcm_data.tobytes()
|
||||
|
||||
# 保存为PCM文件
|
||||
filename = f"tts_output_{req_id[:8]}.pcm"
|
||||
with open(filename, "wb") as f:
|
||||
f.write(pcm_bytes)
|
||||
|
||||
print(f"💾 PCM文件已保存: {filename}, 大小: {len(pcm_bytes)}字节")
|
||||
|
||||
return filename
|
||||
|
||||
elif status == "error":
|
||||
error_msg = result.get("message")
|
||||
print(f"❌ TTS处理失败 [{req_id[:8]}]: {error_msg}")
|
||||
|
||||
return None
|
||||
# 绑定外部回调
|
||||
tts_manager.set_result_callback(handle_tts_result)
|
||||
|
||||
# 3. 设置是否播放(可选,默认True)
|
||||
tts_manager.set_playback_enabled(True) # 设置为False则不播放
|
||||
# 3. 设置是否启用实时播放(默认True)
|
||||
tts_manager.set_playback_enabled(True)
|
||||
|
||||
# 4. 初始化连接
|
||||
await tts_manager.initialize()
|
||||
|
||||
# 5. 异步合成多个文本(非阻塞)
|
||||
# 5. 异步提交多个TTS请求(非阻塞)
|
||||
texts = [
|
||||
"你好,这是第一个排队的TTS请求。",
|
||||
"我是第二个请求,会等第一个处理完再执行。",
|
||||
"第三个请求,支持流式播放和队列管理。",
|
||||
"第四个请求,测试队列的自动消费功能。",
|
||||
"最后一个请求,处理完成后会自动结束。"
|
||||
"你好,这是字节跳动TTS的流式合成测试。",
|
||||
"我支持异步非阻塞调用,多个请求可以排队处理。",
|
||||
"每个请求都会返回唯一的ID,方便你跟踪结果。",
|
||||
"音频数据会通过回调函数返回,支持实时播放和保存文件。",
|
||||
"最后一个测试句子,演示队列的自动消费功能。"
|
||||
]
|
||||
|
||||
req_ids = []
|
||||
for i, text in enumerate(texts):
|
||||
# 提交请求(非阻塞,立即返回)
|
||||
req_id = await tts_manager.synthesize(
|
||||
text,
|
||||
mode="预训练音色",
|
||||
sft_spk="中文女",
|
||||
voice_type=DEFAULT_VOICE_TYPE,
|
||||
encoding=DEFAULT_ENCODING,
|
||||
speed=1.0
|
||||
)
|
||||
req_ids.append(req_id)
|
||||
print(f"已提交请求 {i + 1}: ID={req_id[:8]}")
|
||||
print(f"📤 已提交请求 {i+1}: ID={req_id[:8]}")
|
||||
|
||||
# 可以立即继续其他操作,不需要等待
|
||||
await asyncio.sleep(0.5) # 模拟其他操作
|
||||
# 模拟其他业务逻辑(无需等待TTS完成)
|
||||
await asyncio.sleep(0.3)
|
||||
|
||||
# 6. 可以在这里做其他事情,TTS会在后台处理
|
||||
|
||||
# 7. 等待所有TTS任务完成(可选)
|
||||
# 6. 等待所有TTS任务完成(可选,根据业务需求决定是否等待)
|
||||
await tts_manager.wait_all_completed()
|
||||
|
||||
# 8. 关闭连接
|
||||
# 7. 关闭连接(程序退出前调用)
|
||||
await tts_manager.shutdown()
|
||||
|
||||
|
||||
# ------------------------------
|
||||
# 异步结果回调示例(高级用法)
|
||||
# ------------------------------
|
||||
async def async_result_callback(req_id: str, result: Dict[str, Any]):
|
||||
"""异步结果回调(支持异步操作,如上传音频到服务器)"""
|
||||
if result["status"] == "completed":
|
||||
print(f"\n⚡ 异步处理结果 [{req_id[:8]}] - 开始上传音频...")
|
||||
# 模拟异步上传操作
|
||||
await asyncio.sleep(0.5)
|
||||
print(f"⚡ 异步处理结果 [{req_id[:8]}] - 音频上传完成")
|
||||
|
||||
|
||||
async def advanced_usage_example():
|
||||
"""高级使用示例:动态控制播放和结果处理"""
|
||||
"""高级使用示例:异步回调 + 禁用播放 + 批量请求"""
|
||||
tts_manager = TTSManager()
|
||||
|
||||
tts_manager = TTSManager(ws_url="ws://10.10.10.202:50000/ws/tts")
|
||||
# 设置异步结果回调
|
||||
tts_manager.set_result_callback(async_result_callback)
|
||||
|
||||
# 自定义结果处理器
|
||||
class ResultProcessor:
|
||||
def __init__(self):
|
||||
self.results = {}
|
||||
|
||||
async def process_result(self, req_id: str, result: Dict[str, Any]):
|
||||
"""异步处理结果"""
|
||||
if result.get("status") == "completed":
|
||||
# 保存音频数据等处理
|
||||
self.results[req_id] = result
|
||||
print(f"处理完成: {req_id[:8]}, 音频长度: {result.get('audio_length')}")
|
||||
|
||||
# 这里可以添加自定义逻辑,如保存到文件、发送到网络等
|
||||
|
||||
processor = ResultProcessor()
|
||||
tts_manager.set_result_callback(processor.process_result)
|
||||
|
||||
# 禁用播放,只获取数据
|
||||
# 禁用实时播放(只获取音频数据)
|
||||
tts_manager.set_playback_enabled(False)
|
||||
|
||||
await tts_manager.initialize()
|
||||
|
||||
# 批量合成
|
||||
# 批量提交请求(并行提交)
|
||||
tasks = []
|
||||
for i in range(3):
|
||||
text = f"这是第{i+1}个测试文本,用于演示高级功能。"
|
||||
task = tts_manager.synthesize(text)
|
||||
text = f"这是第{i+1}个高级测试文本,使用异步回调处理结果。"
|
||||
task = tts_manager.synthesize(text, speed=0.9)
|
||||
tasks.append(task)
|
||||
|
||||
# 并行提交所有请求
|
||||
req_ids = await asyncio.gather(*tasks)
|
||||
# 并行提交所有请求
|
||||
req_ids = await asyncio.gather(*tasks)
|
||||
print(f"\n已并行提交 {len(req_ids)} 个请求")
|
||||
|
||||
# 等待完成
|
||||
await tts_manager.wait_all_completed()
|
||||
await tts_manager.shutdown()
|
||||
# 等待所有任务完成
|
||||
await tts_manager.wait_all_completed()
|
||||
await tts_manager.shutdown()
|
||||
|
||||
print(f"所有任务完成,共处理 {len(processor.results)} 个结果")
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 运行示例
|
||||
try:
|
||||
# 基础使用示例
|
||||
# 运行基础使用示例
|
||||
asyncio.run(external_usage_example())
|
||||
|
||||
# 高级使用示例(取消注释运行)
|
||||
# 运行高级使用示例(取消注释)
|
||||
# asyncio.run(advanced_usage_example())
|
||||
|
||||
except KeyboardInterrupt:
|
||||
|
||||
Reference in New Issue
Block a user